migration: establish exact preserved app baseline

This commit is contained in:
leefer
2026-07-30 23:51:48 +08:00
commit e4a9b2e647
389 changed files with 126625 additions and 0 deletions
File diff suppressed because it is too large Load Diff
+215
View File
@@ -0,0 +1,215 @@
from __future__ import annotations
import tempfile
import unittest
import sqlite3
from pathlib import Path
from database import ReviewDatabase
from security import hash_password, verify_password
class AccountAccessTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.database = ReviewDatabase(Path(self.temp.name) / "review.db")
def tearDown(self):
self.temp.cleanup()
def test_first_user_is_admin_and_following_users_are_regular(self):
first = self.database.create_user("admin_user", "salt", "hash")
second = self.database.create_user("member_user", "salt", "hash")
self.assertEqual(first["role"], "admin")
self.assertEqual(second["role"], "user")
self.assertEqual(self.database.user_access(first["id"])["role"], "admin")
self.assertEqual(self.database.user_access(second["id"])["role"], "user")
def test_admin_role_can_also_hold_an_explicit_membership(self):
admin = self.database.create_user("admin_member", "salt", "hash")
updated = self.database.update_membership(
admin["id"],
"active",
"内部会员",
"2026-07-22T00:00:00+00:00",
"2026-08-23T00:00:00+00:00",
)
access = self.database.user_access(admin["id"])
self.assertTrue(updated)
self.assertEqual(access["role"], "admin")
self.assertEqual(access["membership_status"], "active")
def test_membership_mode_system_settings_and_usage_are_persistent(self):
user = self.database.create_user("member_user", "salt", "hash")
self.database.update_user_llm_mode(user["id"], "platform")
updated = self.database.update_membership(
user["id"],
"active",
"内部会员",
"2026-07-22T00:00:00+00:00",
"2026-08-23T00:00:00+00:00",
)
self.database.save_system_setting("credentials", "encrypted")
self.database.record_llm_usage(
user["id"], "mentor", "platform", "model", "success", 1200
)
access = self.database.user_access(user["id"])
self.assertTrue(updated)
self.assertEqual(access["llm_mode"], "platform")
self.assertEqual(access["membership_status"], "active")
self.assertEqual(access["membership_plan"], "内部会员")
self.assertEqual(self.database.get_system_setting("credentials"), "encrypted")
self.assertEqual(
self.database.count_llm_usage_since(
user["id"], "platform", "2026-01-01T00:00:00+00:00"
),
1,
)
def test_password_can_be_rotated_without_changing_account_access(self):
old_salt, old_hash = hash_password("OldPassword123")
user = self.database.create_user("password_user", old_salt, old_hash)
new_salt, new_hash = hash_password("NewPassword456")
self.assertTrue(
self.database.update_user_password(user["id"], new_salt, new_hash)
)
stored = self.database.user_password(user["id"])
self.assertFalse(
verify_password("OldPassword123", stored["password_salt"], stored["password_hash"])
)
self.assertTrue(
verify_password("NewPassword456", stored["password_salt"], stored["password_hash"])
)
self.assertEqual(self.database.user_access(user["id"])["role"], "admin")
def test_latest_real_snapshot_skips_demo_and_supports_strict_previous_date(self):
self.database.save_snapshot(
"20260720", "tushare", {"meta": {"trade_date": "2026-07-20", "source": "tushare"}}
)
self.database.save_snapshot(
"20260721", "demo", {"meta": {"trade_date": "2026-07-21", "source": "demo"}}
)
latest = self.database.get_latest_real_snapshot("20260722")
previous = self.database.get_latest_real_snapshot("20260721", strictly_before=True)
self.assertEqual(latest["meta"]["trade_date"], "2026-07-20")
self.assertEqual(previous["meta"]["trade_date"], "2026-07-20")
def test_stock_master_search_supports_exact_name_and_code(self):
self.database.upsert_stock_master(
[
{
"ts_code": "002141.SZ",
"name": "贤丰控股",
"industry": "元件",
"market": "主板",
"list_date": "20071228",
}
]
)
self.assertEqual(self.database.search_stock_master("贤丰控股")[0]["code"], "002141")
self.assertEqual(self.database.search_stock_master("002141")[0]["name"], "贤丰控股")
def test_review_notes_are_scoped_to_their_owner(self):
first = self.database.create_user("note_owner", "salt", "hash")
second = self.database.create_user("other_reader", "salt", "hash")
note_id = self.database.save_note(
first["id"], "002141", "贤丰控股", "20260721", "只属于甲", "明日观察",
summary="市场缩量修复",
)
first_notes = self.database.list_notes(first["id"], code="002141")
self.assertEqual(len(first_notes), 1)
self.assertEqual(first_notes[0]["summary"], "市场缩量修复")
self.assertEqual(self.database.list_notes(second["id"], code="002141"), [])
with self.assertRaises(ValueError):
self.database.save_note(
second["id"],
"002141",
"贤丰控股",
"20260721",
"越权修改",
"",
note_id,
)
self.assertFalse(self.database.delete_note(second["id"], note_id))
self.assertTrue(self.database.delete_note(first["id"], note_id))
def test_watchlist_is_scoped_to_its_owner(self):
first = self.database.create_user("watch_owner", "salt", "hash")
second = self.database.create_user("other_watcher", "salt", "hash")
self.database.save_watchlist(
first["id"], "002141", "贤丰控股", "元件", "red", "观察承接"
)
self.database.save_watchlist(second["id"], "002141", "贤丰控股", "元件", "blue")
self.assertEqual(self.database.list_watchlist(first["id"])[0]["color"], "red")
self.assertEqual(self.database.list_watchlist(first["id"])[0]["remark"], "观察承接")
self.assertEqual(self.database.list_watchlist(second["id"])[0]["color"], "blue")
self.assertFalse(self.database.delete_watchlist(second["id"], "000001"))
self.assertTrue(self.database.delete_watchlist(first["id"], "002141"))
self.assertEqual(self.database.list_watchlist(first["id"]), [])
self.assertEqual(len(self.database.list_watchlist(second["id"])), 1)
def test_legacy_review_notes_are_assigned_to_first_account(self):
legacy_path = Path(self.temp.name) / "legacy.db"
connection = sqlite3.connect(legacy_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
(username, password_salt, password_hash, created_at, updated_at)
VALUES ('legacy_admin', 'salt', 'hash', '2026-01-01', '2026-01-01');
CREATE TABLE review_notes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code TEXT NOT NULL DEFAULT '',
stock_name TEXT NOT NULL DEFAULT '',
trade_date TEXT NOT NULL,
content TEXT NOT NULL DEFAULT '',
plan TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
INSERT INTO review_notes
(code, stock_name, trade_date, content, plan, created_at, updated_at)
VALUES ('', '', '20260721', '旧复盘', '', '2026-07-21', '2026-07-21');
CREATE TABLE watchlist (
code TEXT PRIMARY KEY,
name TEXT NOT NULL,
sector TEXT NOT NULL DEFAULT '',
color TEXT NOT NULL DEFAULT 'red',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
INSERT INTO watchlist
(code, name, sector, color, created_at, updated_at)
VALUES ('002141', '贤丰控股', '元件', 'red', '2026-07-21', '2026-07-21');
"""
)
connection.commit()
finally:
connection.close()
migrated = ReviewDatabase(legacy_path)
notes = migrated.list_notes(1)
self.assertEqual(len(notes), 1)
self.assertEqual(notes[0]["content"], "旧复盘")
self.assertEqual(migrated.list_watchlist(1)[0]["code"], "002141")
if __name__ == "__main__":
unittest.main()
+302
View File
@@ -0,0 +1,302 @@
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()
+96
View File
@@ -0,0 +1,96 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from alert_service import AlertService
from database import ReviewDatabase
class AlertServiceTests(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("alert_owner", "salt", "hash")
self.other = self.database.create_user("alert_other", "salt", "hash")
self.service = AlertService(self.database)
def tearDown(self) -> None:
self.temp.cleanup()
def test_future_manual_alert_is_visible_but_not_unread_until_due(self):
alert_id = self.service.create_manual(
self.owner["id"],
{
"title": "复核承接",
"content": "开盘不及预期则退出观察",
"code": "002141",
"remind_date": "2026-07-25",
},
)
before = self.service.list_alerts(self.owner["id"], "all", "2026-07-22")
due = self.service.list_alerts(self.owner["id"], "unread", "2026-07-25")
self.assertEqual(before["items"][0]["id"], alert_id)
self.assertFalse(before["items"][0]["due"])
self.assertEqual(before["unread_count"], 0)
self.assertEqual(due["unread_count"], 1)
self.assertTrue(due["items"][0]["due"])
def test_alert_reads_and_mutations_are_scoped_to_owner(self):
alert_id = self.database.save_alert(
self.owner["id"], "manual", "甲的提醒", "", "20260722", "", "owner-only"
)
self.assertEqual(
self.service.list_alerts(self.other["id"], "all", "2026-07-22")["items"], []
)
self.assertFalse(self.database.mark_alert_read(self.other["id"], alert_id))
self.assertFalse(self.database.delete_alert(self.other["id"], alert_id))
self.assertTrue(self.database.mark_alert_read(self.owner["id"], alert_id))
self.assertEqual(
self.service.list_alerts(self.owner["id"], "all", "2026-07-22")["unread_count"], 0
)
self.assertTrue(self.database.delete_alert(self.owner["id"], alert_id))
def test_strategy_alerts_are_idempotent(self):
tracking = {
"batches": [
{
"run_id": 9,
"strategy_name": "修复策略",
"items": [{"code": "002141"}, {"code": "600000"}],
"summary": {
"observed": 2,
"completed": 2,
"t1_win_rate": 50.0,
"average_t5": 3.25,
},
}
]
}
self.service.sync_strategy_tracking(self.owner["id"], tracking)
self.service.sync_strategy_tracking(self.owner["id"], tracking)
alerts = self.database.list_alerts(
self.owner["id"], "99991231", unread_only=False
)
self.assertEqual(len(alerts), 2)
self.assertEqual({item["kind"] for item in alerts}, {"strategy_t1", "strategy_t5"})
def test_mark_all_only_changes_due_alerts(self):
self.database.save_alert(
self.owner["id"], "manual", "今日", "", "20260722", "", "due"
)
self.database.save_alert(
self.owner["id"], "manual", "未来", "", "20260723", "", "future"
)
self.assertEqual(
self.database.mark_all_alerts_read(self.owner["id"], "20260722"), 1
)
future = self.service.list_alerts(self.owner["id"], "unread", "2026-07-23")
self.assertEqual([item["title"] for item in future["items"]], ["未来"])
if __name__ == "__main__":
unittest.main()
+60
View File
@@ -0,0 +1,60 @@
from __future__ import annotations
import unittest
from api_access import required_role
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",
("GET", "/api/heaven/readings"): "member",
("GET", "/api/assistant/messages"): "member",
("POST", "/api/screener/run"): "member",
("POST", "/api/screener/tracking"): "member",
("POST", "/api/screener/tracking/refresh"): "member",
("POST", "/api/mentors/chat"): "member",
("POST", "/api/mentors/preferences"): "member",
("POST", "/api/heaven/interpret"): "member",
("POST", "/api/assistant/chat"): "member",
("DELETE", "/api/screener/strategies/42"): "member",
("DELETE", "/api/screener/tracking/42"): "member",
("DELETE", "/api/mentors/messages"): "member",
("DELETE", "/api/assistant/messages"): "member",
("DELETE", "/api/heaven/readings/42"): "member",
}
for (method, path), role in cases.items():
with self.subTest(method=method, path=path):
self.assertEqual(required_role(method, path), role)
def test_shared_knowledge_mutations_require_admin(self):
cases = (
("POST", "/api/reasons"),
("POST", "/api/seat-aliases"),
("POST", "/api/heaven/sector-phases"),
("DELETE", "/api/heaven/sector-phases/油气开采"),
("POST", "/api/backfill"),
("GET", "/api/admin/settings"),
)
for method, path in cases:
with self.subTest(method=method, path=path):
self.assertEqual(required_role(method, path), "admin")
def test_personal_market_data_routes_need_login_only(self):
cases = (
("GET", "/api/dashboard"),
("GET", "/api/watchlist"),
("POST", "/api/notes"),
("DELETE", "/api/notes/3"),
)
for method, path in cases:
with self.subTest(method=method, path=path):
self.assertEqual(required_role(method, path), "authenticated")
if __name__ == "__main__":
unittest.main()
+53
View File
@@ -0,0 +1,53 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from backend.bootstrap import build_application_container
from backend.bootstrap.settings import environment_credentials
from database import ReviewDatabase
class BootstrapContainerTests(unittest.TestCase):
def test_environment_credentials_preserve_legacy_model_fallbacks(self) -> None:
result = environment_credentials(
{
"TUSHARE_TOKEN": " tushare ",
"IFIND_REFRESH_TOKEN": " refresh ",
"LLM_API_KEY": "legacy-key",
"LLM_BASE_URL": "https://legacy.example/v1",
"LLM_MODEL": "legacy-model",
}
)
self.assertEqual(result["tushare_token"], "tushare")
self.assertEqual(result["ifind_refresh_token"], "refresh")
self.assertEqual(result["platform_llm_primary_api_key"], "legacy-key")
self.assertEqual(result["platform_llm_primary_base_url"], "https://legacy.example/v1")
self.assertEqual(result["platform_llm_primary_model"], "legacy-model")
def test_container_shares_one_database_and_one_ifind_client(self) -> None:
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
public_skills = root / "public"
private_skills = root / "private"
public_skills.mkdir()
private_skills.mkdir()
database = ReviewDatabase(root / "review.db")
container = build_application_container(
database,
{"ifind_refresh_token": "refresh-token", "ifind_access_token": "access-token"},
public_skills,
private_skills,
)
self.assertIs(container.database, database)
self.assertIs(container.screener.database, database)
self.assertIs(container.strategy_tracking.repository.database, database)
self.assertIs(container.alert_service.repository.database, database)
self.assertIs(container.trade_journal.repository.database, database)
self.assertIs(container.chart_data.ifind, container.ifind)
self.assertTrue(container.ifind.configured)
if __name__ == "__main__":
unittest.main()
+133
View File
@@ -0,0 +1,133 @@
from __future__ import annotations
import unittest
from chart_data_provider import ChartDataError, EastmoneyChartClient
from server import DashboardService
class FakeChartClient(EastmoneyChartClient):
def __init__(self) -> None:
super().__init__(cache_ttl_seconds=20)
self.requests: list[tuple[str, dict[str, str]]] = []
def _request_json(self, url, params, referer):
self.requests.append((url, params))
if "trends2" in url:
return {
"data": {
"code": params["secid"].split(".", 1)[1],
"name": "测试行情",
"preClose": 10.0,
"trends": [
"2026-07-24 09:30,10.10,10.20,10.30,10.00,100,1020.00,10.200",
"2026-07-24 09:31,10.20,10.15,10.25,10.10,80,812.00,10.178",
],
}
}
return {
"data": {
"diff": [
{"f12": "BK0474", "f14": "保险Ⅱ"},
{"f12": "BK1040", "f14": "中药Ⅱ"},
]
}
}
class ChartDataProviderTests(unittest.TestCase):
def setUp(self) -> None:
EastmoneyChartClient._cache.clear()
EastmoneyChartClient._board_catalog.clear()
EastmoneyChartClient._board_catalog_at = 0
self.client = FakeChartClient()
def test_stock_intraday_maps_market_and_parses_points(self):
payload = self.client.stock_intraday("601318")
self.assertEqual(self.client.requests[0][1]["secid"], "1.601318")
self.assertEqual(payload["trade_date"], "2026-07-24")
self.assertEqual(payload["points"][0]["time"], "09:30")
self.assertEqual(payload["points"][0]["average"], 10.2)
def test_short_cache_avoids_duplicate_hover_requests(self):
self.client.stock_intraday("002141")
self.client.stock_intraday("002141")
trend_requests = [item for item in self.client.requests if "trends2" in item[0]]
self.assertEqual(len(trend_requests), 1)
def test_index_and_board_use_the_same_chart_shape(self):
index = self.client.index_intraday("000001.SH")
board = self.client.board_intraday("BK0474")
self.assertEqual(index["points"][1]["close"], 10.15)
self.assertEqual(board["points"][1]["volume"], 80.0)
secids = [params["secid"] for url, params in self.client.requests if "trends2" in url]
self.assertIn("1.000001", secids)
self.assertIn("90.BK0474", secids)
def test_invalid_identifier_is_rejected(self):
with self.assertRaises(ChartDataError):
self.client.stock_intraday("abc")
class ChartServiceStub:
@staticmethod
def _payload(code: str, name: str):
return {
"code": code,
"name": name,
"trade_date": "2026-07-24",
"previous_close": 10,
"points": [{"date": "2026-07-24", "time": "09:30", "close": 10.1}],
}
def stock_intraday(self, code):
return self._payload(code, "测试股票")
def index_intraday(self, identifier):
return self._payload(identifier, "上证指数")
def board_intraday(self, identifier, name=""):
return self._payload("BK0474", name)
class ChartDirectoryStub:
@staticmethod
def get_data_snapshot(kind, cache_key):
if (kind, cache_key) != ("search_directory", "ths"):
return None
return {
"schema_version": 2,
"items": [
{"id": "881107.TI", "name": "保险", "type": "sector"},
{"id": "885728.TI", "name": "人工智能", "type": "theme"},
],
}
class IntradayChartServiceTests(unittest.TestCase):
def setUp(self):
self.service = DashboardService.__new__(DashboardService)
self.service.chart_data = ChartServiceStub()
self.service.database = ChartDirectoryStub()
def test_stock_index_sector_and_theme_share_display_only_contract(self):
cases = (
("stock", "601318"),
("index", "000001.SH"),
("sector", "881107.TI"),
("theme", "885728.TI"),
)
for entity_type, identifier in cases:
with self.subTest(entity_type=entity_type):
payload = self.service.get_intraday_chart(entity_type, identifier)
self.assertEqual(payload["entity"]["type"], entity_type)
self.assertEqual(payload["meta"]["trade_date"], "2026-07-24")
self.assertEqual(len(payload["points"]), 1)
self.assertNotIn("source", payload["meta"])
if __name__ == "__main__":
unittest.main()
+82
View File
@@ -0,0 +1,82 @@
from __future__ import annotations
import re
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
STATIC = ROOT / "static"
TOKENS = STATIC / "shared" / "tokens.css"
LEGACY_STYLESHEETS = (
"styles.css",
"renovation.css",
"redesign-v2.css",
"design-system.css",
"theme.css",
)
class CssGovernanceTests(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.html = (STATIC / "index.html").read_text(encoding="utf-8")
cls.tokens = TOKENS.read_text(encoding="utf-8")
def test_token_layer_loads_before_application_styles(self) -> None:
expected_order = (
"/shared/tokens.css",
"/styles.css",
"/renovation.css",
"/redesign-v2.css",
"/design-system.css",
"/theme.css",
"/wentian-v2.css",
)
positions = [self.html.index(path) for path in expected_order]
self.assertEqual(positions, sorted(positions))
def test_token_file_has_three_layer_contract(self) -> None:
for heading in (
"/* Primitive tokens */",
"/* Semantic tokens */",
"/* Component tokens */",
"/* Compatibility aliases.",
):
self.assertIn(heading, self.tokens)
for variable in (
"--color-action:",
"--surface-canvas:",
"--text-primary:",
"--card-bg:",
"--control-height:",
"--sidebar-width:",
):
self.assertIn(variable, self.tokens)
def test_light_and_dark_semantics_share_one_owner(self) -> None:
self.assertIn(':root[data-theme="dark"] {', self.tokens)
global_root = re.compile(r'(?m)^:root(?:\[data-theme="dark"\])?\s*\{')
for filename in LEGACY_STYLESHEETS:
stylesheet = (STATIC / filename).read_text(encoding="utf-8")
self.assertIsNone(global_root.search(stylesheet), filename)
def test_wentian_tokens_remain_isolated(self) -> None:
self.assertNotRegex(self.tokens, r"--wt-[a-z0-9-]+\s*:")
wentian = (STATIC / "wentian-v2.css").read_text(encoding="utf-8")
self.assertRegex(wentian, r"--wt-[a-z0-9-]+\s*:")
def test_compatibility_aliases_cover_historical_layers(self) -> None:
for variable in (
"--blue:",
"--up:",
"--xb-blue-500:",
"--r2-blue:",
"--chart-background:",
"--dragon-profile-list-width:",
):
self.assertIn(variable, self.tokens)
if __name__ == "__main__":
unittest.main()
+542
View File
@@ -0,0 +1,542 @@
import sqlite3
import tempfile
import unittest
from datetime import datetime, timedelta
from pathlib import Path
from database import ReviewDatabase
from screener import (
ADVANCED_CURATED_STRATEGIES,
CURATED_STRATEGIES,
FACTOR_FIELDS,
FACTOR_GROUPS,
ScreenerEngine,
_broken_reversal_metrics,
_earnings_event_rows,
_popularity_factor_rows,
_risk_flags,
_rsi,
_quarter_periods,
)
from server import DashboardService, automatic_screener_jobs
class CuratedScreenerTests(unittest.TestCase):
def test_curated_library_contains_original_and_advanced_strategies(self):
self.assertEqual(19, len(ADVANCED_CURATED_STRATEGIES))
self.assertEqual(29, len(CURATED_STRATEGIES))
self.assertEqual(29, len({item["name"] for item in CURATED_STRATEGIES}))
self.assertTrue(
{"行业动量轮动", "主力资金行业流入"}.issubset(
{item["name"] for item in CURATED_STRATEGIES}
)
)
self.assertTrue(
all(item["formula"]["meta"]["library"] == "curated" for item in CURATED_STRATEGIES)
)
self.assertTrue(
{
"景气-趋势-拥挤三维行业打分",
"大小盘/成长价值风格切换(元策略)",
"业绩超预期漂移(SUE/PEAD)",
"多因子综合打分(IC动态加权)",
"热度突增潜伏(另类数据)",
"机构榜溢价",
}.issubset({item["name"] for item in CURATED_STRATEGIES})
)
def test_every_curated_strategy_explains_environment_and_failure_risk(self):
for strategy in CURATED_STRATEGIES:
meta = strategy["formula"]["meta"]
self.assertTrue(meta.get("suitable_environment"), strategy["name"])
self.assertTrue(meta.get("failure_risk"), strategy["name"])
self.assertNotIn("emotion_gate", meta, strategy["name"])
def test_automatic_curated_jobs_are_not_filtered_by_market_regime(self):
strategies = [
{
"name": "阶段策略",
"regimes": ["retreat"],
"formula": {"meta": {"library": "stage"}},
},
*CURATED_STRATEGIES,
]
for regime in ("ice", "repair", "fermentation", "climax", "divergence", "retreat"):
jobs = automatic_screener_jobs(strategies, regime)
curated_names = {
job["strategy"]["name"] for job in jobs if job["mode"] == "curated"
}
self.assertEqual(
{strategy["name"] for strategy in CURATED_STRATEGIES},
curated_names,
regime,
)
def test_curated_risk_flags_do_not_reintroduce_regime_gating(self):
row = {
"pct_chg": 0,
"return_10d": 0,
"volatility_10d": 0,
"amount_billion": 5,
}
self.assertIn("市场处于退潮阶段,策略可能选择空仓", _risk_flags(row, "retreat"))
self.assertNotIn(
"市场处于退潮阶段,策略可能选择空仓",
_risk_flags(row, "retreat", include_regime_risk=False),
)
def test_every_curated_formula_uses_supported_factors(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
engine = ScreenerEngine(database)
for strategy in CURATED_STRATEGIES:
formula = engine.validate_formula(strategy["formula"])
fields = {
item["field"]
for item in formula["filters"] + formula["score"]
}
self.assertTrue(fields.issubset(FACTOR_FIELDS), strategy["name"])
def test_server_gate_blocks_specialized_strategies_until_sources_are_ready(self):
factor_dates = [f"2026{index + 1:04d}" for index in range(260)]
health = {
"market": True,
"auction": True,
"benchmark": True,
"valuation": True,
"fundamental": True,
"dividend_history": True,
"moneyflow_history": True,
"earnings_events": False,
"popularity": False,
"institutions": False,
}
expected = {
"业绩超预期漂移(SUE/PEAD)": "业绩预告与快报",
"热度突增潜伏(另类数据)": "当日人气榜",
"机构榜溢价": "龙虎榜机构席位",
}
by_name = {strategy["name"]: strategy for strategy in CURATED_STRATEGIES}
for name, missing_label in expected.items():
self.assertEqual(
[missing_label],
DashboardService._strategy_missing_data(
by_name[name], factor_dates, health
),
name,
)
ready_health = {
**health,
"earnings_events": True,
"popularity": True,
"institutions": True,
}
for name in expected:
self.assertEqual(
[],
DashboardService._strategy_missing_data(
by_name[name], factor_dates, ready_health
),
name,
)
def test_factor_groups_cover_every_quant_factor(self):
grouped = [field for fields in FACTOR_GROUPS.values() for field in fields]
self.assertEqual(set(FACTOR_FIELDS), set(grouped))
self.assertEqual(len(grouped), len(set(grouped)))
def test_database_migrates_valuation_and_fundamental_columns(self):
with tempfile.TemporaryDirectory() as root:
path = Path(root) / "review.db"
ReviewDatabase(path)
connection = sqlite3.connect(path)
try:
indicator_columns = {
row[1] for row in connection.execute("PRAGMA table_info(daily_indicators)")
}
tables = {
row[0] for row in connection.execute(
"SELECT name FROM sqlite_master WHERE type='table'"
)
}
finally:
connection.close()
self.assertTrue({"pe_ttm", "pb", "ps_ttm", "dv_ttm"}.issubset(indicator_columns))
self.assertIn("fundamental_indicators", tables)
self.assertIn("benchmark_bars", tables)
self.assertIn("earnings_events", tables)
self.assertIn("popularity_factors", tables)
self.assertIn("lhb_institution_daily", tables)
def test_advanced_strategies_declare_history_and_backtest_contracts(self):
for strategy in ADVANCED_CURATED_STRATEGIES:
meta = strategy["formula"]["meta"]
self.assertGreaterEqual(meta["history_days"], 80, strategy["name"])
self.assertGreaterEqual(meta["backtest_days"], 1, strategy["name"])
self.assertGreater(meta["take_profit"], 0, strategy["name"])
self.assertLess(meta["stop_loss"], 0, strategy["name"])
def test_quarter_periods_stop_at_selected_date(self):
periods = _quarter_periods("20260722", 5)
self.assertEqual(
["20250630", "20250930", "20251231", "20260331", "20260630"],
periods,
)
def test_factor_health_summary_uses_availability_counts(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
with database.connect() as connection:
connection.execute(
"INSERT INTO daily_bars (trade_date, ts_code) VALUES (?, ?)",
("20260722", "600000.SH"),
)
connection.execute(
"""
INSERT INTO daily_indicators
(trade_date, ts_code, pe_ttm)
VALUES (?, ?, ?)
""",
("20260722", "600000.SH", 8.5),
)
connection.executemany(
"INSERT INTO daily_indicators (trade_date, ts_code) VALUES (?, ?)",
[(f"{year}1231", f"{year % 100:02d}0000.SZ") for year in range(2022, 2026)],
)
connection.execute(
"INSERT INTO auction_factors (trade_date, ts_code) VALUES (?, ?)",
("20260722", "600000.SH"),
)
connection.executemany(
"""
INSERT INTO fundamental_indicators (end_date, ann_date, ts_code, roe)
VALUES (?, ?, ?, ?)
""",
[
("20251231", "20260430", f"{index:06d}.SZ", 10.0)
for index in range(100)
],
)
connection.executemany(
"INSERT INTO benchmark_bars (trade_date, ts_code, close) VALUES (?, ?, ?)",
[(f"2026{index + 1:04d}", "000300.SH", 4000 + index) for index in range(60)],
)
health = database.factor_health_summary("20260722")
self.assertTrue(health["market"])
self.assertTrue(health["auction"])
self.assertTrue(health["valuation"])
self.assertTrue(health["fundamental"])
self.assertTrue(health["dividend_history"])
self.assertTrue(health["benchmark"])
self.assertEqual(health["valuation_rows"], 1)
self.assertEqual(health["fundamental_rows"], 100)
self.assertEqual(health["dividend_years"], 5)
def test_moneyflow_health_requires_the_latest_five_market_dates(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
dates = [f"202607{day:02d}" for day in range(20, 25)]
database.upsert_daily_bars([
{
"trade_date": trade_date, "ts_code": "600000.SH",
"open": 10, "high": 10.2, "low": 9.8, "close": 10,
"pct_chg": 0, "vol": 1000, "amount": 100000,
}
for trade_date in dates
])
database.upsert_moneyflow([
{"trade_date": "20260105", "ts_code": "600000.SH", "net_mf_amount": 10}
] * 5)
self.assertFalse(database.factor_health_summary(dates[-1])["moneyflow_history"])
database.upsert_moneyflow([
{"trade_date": trade_date, "ts_code": "600000.SH", "net_mf_amount": 10}
for trade_date in dates
])
health = database.factor_health_summary(dates[-1])
self.assertTrue(health["moneyflow_history"])
self.assertEqual(health["moneyflow_dates"], 5)
def test_technical_helpers_detect_rsi_and_daily_reversal_path(self):
self.assertLess(_rsi([10, 9, 8, 7, 6, 5, 4], 6), 1)
rows = [
{"close": 10, "high": 10, "vol": 100},
{"close": 11, "high": 11, "vol": 120},
{"close": 12, "high": 12, "vol": 130},
{"close": 11.2, "high": 11.8, "vol": 100},
{"close": 12.5, "high": 12.5, "vol": 140},
]
metrics = _broken_reversal_metrics(
rows, [False, True, True, False, True], "600000", "示例"
)
self.assertEqual(metrics["signal"], 1)
self.assertEqual(metrics["days"], 1)
def test_factor_builder_generates_long_window_and_benchmark_factors(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
database.upsert_stock_master([
{
"ts_code": "600000.SH", "name": "趋势样本", "industry": "银行",
"market": "主板", "list_date": "20000101",
}
])
dates = []
cursor = datetime(2025, 6, 1)
while len(dates) < 260:
if cursor.weekday() < 5:
dates.append(cursor.strftime("%Y%m%d"))
cursor += timedelta(days=1)
bars = []
benchmarks = []
indicators = []
for index, trade_date in enumerate(dates):
close = 10 + index * 0.05
bars.append({
"trade_date": trade_date, "ts_code": "600000.SH",
"open": close - 0.02, "high": close + 0.08, "low": close - 0.08,
"close": close, "pct_chg": 0.25, "vol": 1000 + index,
"amount": 200000,
})
benchmarks.append({
"trade_date": trade_date, "ts_code": "000300.SH",
"close": 4000 + index, "pct_chg": 0.02,
})
if index >= 250:
indicators.append({
"trade_date": trade_date, "ts_code": "600000.SH",
"turnover_rate": 2, "volume_ratio": 1,
})
database.upsert_daily_bars(bars)
database.upsert_benchmark_bars(benchmarks)
database.upsert_daily_indicators(indicators)
factors, actual_date = ScreenerEngine(database).build_factors(
dates[-1], history_days=260
)
self.assertEqual(actual_date, dates[-1])
self.assertEqual(len(factors), 1)
factor = factors[0]
self.assertEqual(factor["ma_bull_alignment"], 1)
self.assertEqual(factor["rs_high_120"], 1)
self.assertGreater(factor["momentum_60_5"], 0)
self.assertEqual(factor["momentum_60_5_rank"], 0)
def test_factor_builder_generates_sector_momentum_and_five_day_flow(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
stocks = [
("600001.SH", "动量样本", "电子", 0.16, 180),
("600002.SH", "对照样本", "银行", 0.02, -40),
]
database.upsert_stock_master([
{
"ts_code": code, "name": name, "industry": industry,
"market": "主板", "list_date": "20000101",
}
for code, name, industry, _, _ in stocks
])
dates = []
cursor = datetime(2026, 4, 1)
while len(dates) < 80:
if cursor.weekday() < 5:
dates.append(cursor.strftime("%Y%m%d"))
cursor += timedelta(days=1)
bars = []
for index, trade_date in enumerate(dates):
for code, _, _, slope, _ in stocks:
close = 10 + index * slope
bars.append({
"trade_date": trade_date, "ts_code": code,
"open": close - 0.03, "high": close + 0.08,
"low": close - 0.08, "close": close,
"pct_chg": slope, "vol": 1000 + index,
"amount": 300000,
})
database.upsert_daily_bars(bars)
database.upsert_daily_indicators([
{
"trade_date": dates[-1], "ts_code": code,
"turnover_rate": 2, "volume_ratio": 1,
"circ_mv": 1000000, "total_mv": 1500000,
}
for code, *_ in stocks
])
database.upsert_moneyflow([
{
"trade_date": trade_date, "ts_code": code,
"net_mf_amount": daily_flow,
}
for trade_date in dates[-5:]
for code, _, _, _, daily_flow in stocks
])
factors, _ = ScreenerEngine(database).build_factors(
dates[-1], history_days=80
)
by_code = {item["ts_code"]: item for item in factors}
leader = by_code["600001.SH"]
laggard = by_code["600002.SH"]
self.assertGreater(leader["return_20d"], laggard["return_20d"])
self.assertEqual(leader["sector_momentum_rank"], 1)
self.assertEqual(laggard["sector_momentum_rank"], 0)
self.assertGreater(leader["net_flow_5d_million"], 0)
self.assertLess(laggard["net_flow_5d_million"], 0)
self.assertEqual(leader["sector_flow_rank"], 1)
def test_stage_three_event_and_composite_factors_are_date_scoped(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
stocks = [
("600001.SH", "成长样本", "电子", 0.08),
("600002.SH", "价值样本", "银行", 0.02),
]
database.upsert_stock_master([
{
"ts_code": code, "name": name, "industry": industry,
"market": "主板", "list_date": "20000101",
}
for code, name, industry, _ in stocks
])
dates = []
cursor = datetime(2026, 3, 1)
while len(dates) < 80:
if cursor.weekday() < 5:
dates.append(cursor.strftime("%Y%m%d"))
cursor += timedelta(days=1)
database.upsert_daily_bars([
{
"trade_date": trade_date, "ts_code": code,
"open": 10 + index * slope - 0.02,
"high": 10 + index * slope + 0.08,
"low": 10 + index * slope - 0.08,
"close": 10 + index * slope,
"pct_chg": slope, "vol": 1000 + index, "amount": 300000,
}
for index, trade_date in enumerate(dates)
for code, _, _, slope in stocks
])
database.upsert_daily_indicators([
{
"trade_date": dates[-1], "ts_code": "600001.SH",
"turnover_rate": 3, "volume_ratio": 1.4, "total_mv": 900000,
"circ_mv": 700000, "pe_ttm": 25, "pb": 3, "ps_ttm": 4,
},
{
"trade_date": dates[-1], "ts_code": "600002.SH",
"turnover_rate": 1, "volume_ratio": 0.9, "total_mv": 5000000,
"circ_mv": 4000000, "pe_ttm": 8, "pb": 0.8, "ps_ttm": 1,
},
])
database.upsert_fundamental_indicators([
{
"end_date": "20260331", "ann_date": dates[-10],
"ts_code": "600001.SH", "roe": 16, "roic": 13,
"grossprofit_margin": 35, "netprofit_yoy": 45, "or_yoy": 30,
},
{
"end_date": "20260331", "ann_date": dates[-10],
"ts_code": "600002.SH", "roe": 9, "roic": 7,
"grossprofit_margin": 18, "netprofit_yoy": 5, "or_yoy": 3,
},
])
database.upsert_earnings_events([{
"end_date": "20260331", "ann_date": dates[-3],
"ts_code": "600001.SH", "forecast_profit": 100,
"actual_profit": 125, "surprise_pct": 25,
"revenue_yoy": 30, "netprofit_yoy": 45,
"source": "forecast+express",
}])
database.upsert_popularity_factors([{
"trade_date": dates[-1], "ts_code": "600001.SH",
"ths_rank": 5, "dc_rank": 8, "combined_score": 75,
"rank_change": 12, "dual_source": True,
}])
database.upsert_lhb_institutions([{
"trade_date": dates[-1], "ts_code": "600001.SH",
"exalter": "机构专用", "buy": 80_000_000,
"sell": 20_000_000, "net_buy": 60_000_000,
}])
factors, actual_date = ScreenerEngine(database).build_factors(
dates[-1], history_days=80
)
by_code = {item["ts_code"]: item for item in factors}
factor = by_code["600001.SH"]
self.assertEqual(actual_date, dates[-1])
self.assertEqual(factor["earnings_days_since_announce"], 2)
self.assertEqual(factor["earnings_surprise_pct"], 25)
self.assertEqual(factor["popularity_score"], 75)
self.assertEqual(factor["popularity_dual_source"], 1)
self.assertEqual(factor["institution_net_buy_million"], 60)
self.assertEqual(factor["institution_seat_count"], 1)
self.assertIsNotNone(factor["sector_composite_score"])
self.assertIsNotNone(factor["style_fit_score"])
self.assertIsNotNone(factor["multi_factor_composite"])
health = database.factor_health_summary(dates[-1])
self.assertTrue(health["earnings_events"])
self.assertTrue(health["popularity"])
self.assertTrue(health["institutions"])
def test_stage_three_sources_normalize_units_and_rank_changes(self):
earnings = _earnings_event_rows(
[{
"ts_code": "600001.SH", "ann_date": "20260401",
"end_date": "20260331", "net_profit_min": 10000,
"net_profit_max": 12000,
}],
[{
"ts_code": "600001.SH", "ann_date": "20260420",
"end_date": "20260331", "n_income": 132_000_000,
"yoy_net_profit": 30, "yoy_sales": 18,
}],
"20260420",
)
self.assertEqual(len(earnings), 1)
self.assertEqual(round(earnings[0]["actual_profit"]), 13200)
self.assertEqual(round(earnings[0]["surprise_pct"]), 20)
popularity = _popularity_factor_rows(
"20260420",
[{"data_type": "热股", "ts_code": "600001.SH", "rank": 5}],
[{"data_type": "A股市场", "ts_code": "600001.SH", "rank": 8}],
[{"data_type": "热股", "ts_code": "600001.SH", "rank": 20}],
[{"data_type": "A股市场", "ts_code": "600001.SH", "rank": 30}],
)
self.assertEqual(len(popularity), 1)
self.assertEqual(popularity[0]["rank_change"], 15)
self.assertTrue(popularity[0]["dual_source"])
def test_screen_reports_signal_health(self):
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
engine = ScreenerEngine(database)
formula = {
"universe": {"exclude_st": True, "listed_days_min": 0},
"filters": [{"field": "pct_chg", "op": ">", "value": 0}],
"score": [{"field": "amount_billion", "weight": 1, "direction": "desc"}],
"limit": 5,
"min_score": 0,
}
result = engine.screen(
0, "20260724", formula, "repair", "健康检查", False,
mode="curated",
prepared_factors=[{
"ts_code": "600000.SH", "code": "600000", "name": "浦发银行",
"sector": "银行", "listed_days": 1000, "pct_chg": 1,
"amount_billion": 5, "price": 10, "return_5d": 1,
"volume_ratio_5d": 1, "sector_strength": 50,
}],
prepared_date="20260724",
)
health = result["meta"]["health"]
self.assertEqual(health["status"], "normal")
self.assertEqual(health["signal_count"], 1)
self.assertEqual(health["coverage"], 100)
if __name__ == "__main__":
unittest.main()
+116
View File
@@ -0,0 +1,116 @@
from __future__ import annotations
import copy
import unittest
from server import DashboardService
class SnapshotDatabase:
def __init__(self, snapshot, latest=None):
self.snapshot = snapshot
self.latest = latest
self.aliases = {}
def get_snapshot(self, _trade_date):
return copy.deepcopy(self.snapshot)
def reason_overrides(self, _trade_date):
return {}
def get_data_snapshot(self, kind, cache_key):
return copy.deepcopy(self.aliases.get((kind, cache_key)))
def save_data_snapshot(self, kind, cache_key, _source, payload):
self.aliases[(kind, cache_key)] = copy.deepcopy(payload)
def save_snapshot(self, _trade_date, _source, payload):
self.snapshot = copy.deepcopy(payload)
def get_latest_real_snapshot(self, _trade_date, strictly_before=False):
return copy.deepcopy(self.latest)
class DashboardCacheTests(unittest.TestCase):
def service(self, snapshot):
service = object.__new__(DashboardService)
service.database = SnapshotDatabase(snapshot)
return service
def test_cached_dashboard_skips_sentiment_rebuild_when_fields_are_complete(self):
snapshot = {
"meta": {"source": "tushare", "trade_date": "2026-07-22"},
"overview": {
"sentiment_score": 32,
"sentiment_label": "weak",
"sentiment_phase": "retreat",
"sentiment_direction": "cooling",
"sentiment_components": {},
"sentiment_engine_version": 2,
},
}
service = self.service(snapshot)
service._enrich_dashboard_sentiment = lambda *_args: self.fail(
"complete cached sentiment must not be rebuilt"
)
payload = service.get_dashboard("2026-07-22")
self.assertTrue(payload["meta"]["cached"])
self.assertEqual(payload["overview"]["sentiment_score"], 32)
def test_cached_dashboard_rebuilds_legacy_snapshot_missing_sentiment(self):
snapshot = {
"meta": {"source": "tushare", "trade_date": "2026-07-22"},
"overview": {"limit_up_count": 20},
}
service = self.service(snapshot)
calls = []
def enrich(payload, trade_date):
calls.append(trade_date)
payload["overview"].update({
"sentiment_score": 20,
"sentiment_label": "weak",
"sentiment_phase": "ice",
"sentiment_direction": "cooling",
"sentiment_components": {},
"sentiment_engine_version": 2,
})
return payload
service._enrich_dashboard_sentiment = enrich
payload = service.get_dashboard("2026-07-22")
self.assertEqual(calls, ["20260722"])
self.assertEqual(payload["overview"]["sentiment_phase"], "ice")
def test_weekend_dashboard_reuses_latest_close_without_external_sync(self):
latest = {
"meta": {"source": "tushare", "trade_date": "2026-07-24"},
"overview": {
"sentiment_score": 32,
"sentiment_label": "weak",
"sentiment_phase": "retreat",
"sentiment_direction": "cooling",
"sentiment_components": {},
},
}
service = object.__new__(DashboardService)
service.database = SnapshotDatabase(None, latest)
service.sync_dashboard = lambda *_args: self.fail(
"weekend refresh must not call the external synchronization path"
)
first = service.get_dashboard("2026-07-25")
service.database.latest = None
second = service.get_dashboard("2026-07-25")
self.assertTrue(first["meta"]["carried_forward"])
self.assertEqual(first["meta"]["trade_date"], "2026-07-24")
self.assertEqual(second["meta"]["requested_date"], "2026-07-25")
if __name__ == "__main__":
unittest.main()
+151
View File
@@ -0,0 +1,151 @@
from __future__ import annotations
import unittest
from datetime import datetime, timedelta
from backend.data import (
DataPolicyError,
DataQualityError,
DataSourcePolicy,
QualityEvidence,
build_data_gateway,
)
from backend.data.quality import market_timezone
class DataGatewayTests(unittest.TestCase):
def test_policy_allows_registered_calculation_source(self) -> None:
policy = DataSourcePolicy.load()
contract = policy.assert_allowed(
"market.stock_daily", "tushare", "calculation"
)
self.assertEqual(contract.primary, "tushare")
def test_policy_rejects_public_web_source_for_calculation(self) -> None:
policy = DataSourcePolicy.load()
with self.assertRaises(DataPolicyError):
policy.assert_allowed(
"observation.realtime_indices", "eastmoney", "calculation"
)
def test_policy_rejects_blocked_dataset(self) -> None:
policy = DataSourcePolicy.load()
with self.assertRaises(DataPolicyError):
policy.assert_allowed("market.level2", "unresolved", "display")
def test_gateway_uses_live_token_supplier_and_shared_ifind(self) -> None:
token = {"value": "first"}
gateway = build_data_gateway(
{"ifind_refresh_token": "refresh", "ifind_access_token": "access"},
lambda: token["value"],
)
self.assertEqual(gateway.tushare().token, "first")
token["value"] = "second"
self.assertEqual(gateway.tushare().token, "second")
self.assertIs(gateway.chart_data.ifind, gateway.ifind)
def test_server_has_no_direct_runtime_tushare_construction(self) -> None:
from pathlib import Path
source = (Path(__file__).resolve().parents[1] / "server.py").read_text(encoding="utf-8")
self.assertEqual(source.count("TushareClient(self.token)"), 1)
self.assertIn("return gateway.tushare()", source)
def test_quality_gate_accepts_matching_daily_evidence(self) -> None:
timezone = market_timezone()
now = datetime(2026, 7, 29, 16, 0, tzinfo=timezone)
gateway = build_data_gateway({})
report = gateway.require_quality(
QualityEvidence(
dataset_id="market.stock_daily",
provider_id="tushare",
data_time="2026-07-29",
observed_at=now,
actual_count=5000,
expected_count=5000,
adjustment="current-unadjusted",
units={
"open": "CNY/share", "high": "CNY/share", "low": "CNY/share",
"close": "CNY/share", "pct_chg": "percent",
"volume_shares": "share", "amount_yuan": "CNY",
},
),
"calculation",
now,
)
self.assertTrue(report.accepted)
self.assertEqual(report.coverage_ratio, 1.0)
def test_quality_gate_rejects_stale_dynamic_auction(self) -> None:
timezone = market_timezone()
now = datetime(2026, 7, 29, 9, 24, tzinfo=timezone)
gateway = build_data_gateway({})
with self.assertRaises(DataQualityError):
gateway.require_quality(
QualityEvidence(
dataset_id="market.auction_dynamic",
provider_id="ifind",
data_time=now - timedelta(seconds=30),
observed_at=now - timedelta(seconds=29),
units={
"price": "CNY/share", "volume_shares": "share",
"amount_yuan": "CNY", "pre_close": "CNY/share",
"turnover_rate_pct": "percent", "volume_ratio": "ratio",
"float_share": "share",
},
),
"calculation",
now,
)
def test_quality_gate_rejects_low_coverage_and_wrong_adjustment(self) -> None:
timezone = market_timezone()
now = datetime(2026, 7, 29, 16, 0, tzinfo=timezone)
gateway = build_data_gateway({})
report = gateway.quality.evaluate(
QualityEvidence(
dataset_id="market.stock_daily",
provider_id="tushare",
data_time="2026-07-29",
observed_at=now,
actual_count=4000,
expected_count=5000,
adjustment="forward1",
),
"calculation",
now,
)
self.assertFalse(report.accepted)
self.assertTrue(any("Coverage" in issue for issue in report.issues))
self.assertTrue(any("Adjustment" in issue for issue in report.issues))
def test_quality_gate_enforces_financial_point_in_time(self) -> None:
timezone = market_timezone()
now = datetime(2026, 7, 29, 16, 0, tzinfo=timezone)
gateway = build_data_gateway({})
report = gateway.quality.evaluate(
QualityEvidence(
dataset_id="market.fundamentals",
provider_id="tushare",
data_time="2026-06-30",
observed_at=now,
available_at="2026-08-15",
),
"calculation",
now,
)
self.assertFalse(report.accepted)
self.assertTrue(any("not available" in issue for issue in report.issues))
def test_provider_chain_never_silently_promotes_display_fallback(self) -> None:
gateway = build_data_gateway({})
self.assertEqual(
gateway.provider_chain("chart.intraday", "display"),
("ifind", "eastmoney"),
)
with self.assertRaises(RuntimeError):
gateway.provider_chain("chart.intraday", "calculation")
if __name__ == "__main__":
unittest.main()
+81
View File
@@ -0,0 +1,81 @@
from __future__ import annotations
import sqlite3
import tempfile
import unittest
from pathlib import Path
from backend.database import Migration, MigrationError, MigrationRunner
from database import ReviewDatabase
class DatabaseMigrationTests(unittest.TestCase):
def test_fresh_database_records_the_adopted_schema_once(self) -> None:
with tempfile.TemporaryDirectory() as root:
path = Path(root) / "review.db"
database = ReviewDatabase(path)
with database.connect() as connection:
rows = connection.execute(
"SELECT version, name FROM schema_migrations"
).fetchall()
self.assertEqual(
[(row["version"], row["name"]) for row in rows],
[
("0001", "adopt_legacy_schema"),
("0002", "create_job_runs"),
("0003", "extend_llm_audit"),
],
)
ReviewDatabase(path)
with database.connect() as connection:
count = connection.execute(
"SELECT COUNT(*) AS count FROM schema_migrations"
).fetchone()["count"]
self.assertEqual(count, 3)
def test_connection_factory_enables_required_pragmas(self) -> None:
with tempfile.TemporaryDirectory() as root:
database = ReviewDatabase(Path(root) / "review.db")
with database.connect() as connection:
self.assertEqual(connection.execute("PRAGMA foreign_keys").fetchone()[0], 1)
self.assertEqual(connection.execute("PRAGMA journal_mode").fetchone()[0], "wal")
self.assertEqual(connection.execute("PRAGMA busy_timeout").fetchone()[0], 20000)
def test_failed_migration_rolls_back_and_is_not_recorded(self) -> None:
connection = sqlite3.connect(":memory:")
self.addCleanup(connection.close)
connection.row_factory = sqlite3.Row
def fail(conn: sqlite3.Connection) -> None:
conn.execute("CREATE TABLE should_rollback (id INTEGER)")
raise RuntimeError("stop")
migration = Migration("9000", "failure", fail, "failure:v1")
with self.assertRaises(MigrationError):
MigrationRunner().apply(connection, (migration,))
tables = {
row["name"]
for row in connection.execute(
"SELECT name FROM sqlite_master WHERE type = 'table'"
)
}
self.assertNotIn("should_rollback", tables)
self.assertEqual(
connection.execute("SELECT COUNT(*) FROM schema_migrations").fetchone()[0],
0,
)
def test_applied_migration_checksum_is_immutable(self) -> None:
connection = sqlite3.connect(":memory:")
self.addCleanup(connection.close)
connection.row_factory = sqlite3.Row
first = Migration("9001", "example", lambda conn: None, "example:v1")
changed = Migration("9001", "example", lambda conn: None, "example:v2")
runner = MigrationRunner()
runner.apply(connection, (first,))
with self.assertRaises(MigrationError):
runner.apply(connection, (changed,))
if __name__ == "__main__":
unittest.main()
+36
View File
@@ -0,0 +1,36 @@
from __future__ import annotations
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
class DeploymentContractTests(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.compose = (ROOT / "compose.yaml").read_text(encoding="utf-8")
cls.dockerfile = (ROOT / "Dockerfile").read_text(encoding="utf-8")
cls.dockerignore = (ROOT / ".dockerignore").read_text(encoding="utf-8")
cls.gitignore = (ROOT / ".gitignore").read_text(encoding="utf-8")
def test_compose_exposes_only_requested_lan_port(self):
self.assertIn('"0.0.0.0:8765:8765/tcp"', self.compose)
self.assertIn("read_only: true", self.compose)
self.assertIn("target: /app/data", self.compose)
self.assertIn("no-new-privileges:true", self.compose)
def test_image_runs_as_non_root_with_healthcheck(self):
self.assertIn("USER xiaobai", self.dockerfile)
self.assertIn("HEALTHCHECK", self.dockerfile)
self.assertIn('"--host", "0.0.0.0", "--port", "8765"', self.dockerfile)
def test_secrets_and_runtime_data_are_not_copied_into_image(self):
for pattern in (".env", "data/private-mentor-skills/", "data/*.db", "data/*.db-wal", "data/*.db-shm"):
self.assertIn(pattern, self.dockerignore)
self.assertIn("data/private-mentor-skills/", self.gitignore)
if __name__ == "__main__":
unittest.main()
+56
View File
@@ -0,0 +1,56 @@
from __future__ import annotations
import ast
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
FEATURES = ROOT / "backend" / "features"
class FeatureBoundaryTests(unittest.TestCase):
def test_feature_services_do_not_import_http_or_provider_adapters(self) -> None:
forbidden = {
"server",
"tushare_client",
"ifind_client",
"chart_data_provider",
"realtime_aggregator",
}
violations = []
for path in FEATURES.rglob("*.py"):
tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
for node in ast.walk(tree):
names = []
if isinstance(node, ast.Import):
names = [alias.name for alias in node.names]
elif isinstance(node, ast.ImportFrom) and node.module:
names = [node.module]
for name in names:
if name.split(".")[0] in forbidden:
violations.append(f"{path.relative_to(ROOT)} -> {name}")
self.assertEqual(violations, [])
def test_legacy_service_modules_are_compatibility_exports_only(self) -> None:
for filename in ("alert_service.py", "trade_journal.py", "strategy_tracking.py"):
tree = ast.parse((ROOT / filename).read_text(encoding="utf-8"))
definitions = [
node for node in tree.body
if isinstance(node, (ast.ClassDef, ast.FunctionDef, ast.AsyncFunctionDef))
]
self.assertEqual(definitions, [], filename)
def test_each_migrated_feature_owns_one_application_service(self) -> None:
expected = {
"alerts/service.py": "AlertService",
"review/trade_journal.py": "TradeJournalService",
"screener/tracking.py": "StrategyTrackingService",
}
for relative, class_name in expected.items():
tree = ast.parse((FEATURES / relative).read_text(encoding="utf-8"))
self.assertIn(class_name, {node.name for node in tree.body if isinstance(node, ast.ClassDef)})
if __name__ == "__main__":
unittest.main()
+89
View File
@@ -0,0 +1,89 @@
from __future__ import annotations
import unittest
from heaven_engine import build_five_phase_field
class FivePhaseFrameworkTests(unittest.TestCase):
def test_public_field_uses_year_current_qi_day_contract(self):
field = build_five_phase_field("2026-06-15")
self.assertEqual(
field["framework"]["weights"],
{
"year_movement": 30,
"sitian_zaiquan": 20,
"sitian": 15,
"zaiquan": 5,
"host_qi": 20,
"guest_qi": 25,
"day": 5,
},
)
self.assertEqual(
[(layer["id"], layer["weight"]) for layer in field["framework"]["layers"]],
[("year", 50), ("current", 45), ("day", 5)],
)
self.assertEqual(sum(item["score"] for item in field["balance"]), 100)
def test_sitian_and_zaiquan_follow_half_year_dominance(self):
first_half = build_five_phase_field("2026-06-15")
second_half = build_five_phase_field("2026-08-20")
self.assertEqual(first_half["framework"]["weights"]["sitian"], 15)
self.assertEqual(first_half["framework"]["weights"]["zaiquan"], 5)
self.assertEqual(first_half["six_qi"]["ruling"], "司天")
self.assertEqual(second_half["framework"]["weights"]["sitian"], 5)
self.assertEqual(second_half["framework"]["weights"]["zaiquan"], 15)
self.assertEqual(second_half["six_qi"]["ruling"], "在泉")
def test_guest_host_relation_and_anchor_alignment_are_explicit(self):
third_qi = build_five_phase_field("2026-06-15")
final_qi = build_five_phase_field("2026-12-10")
controlled = build_five_phase_field("2025-02-10")
self.assertEqual(third_qi["framework"]["relations"]["guest_host"]["label"], "客主同气")
self.assertEqual(third_qi["six_qi"]["alignment"], "司天同位")
self.assertEqual(final_qi["framework"]["relations"]["guest_host"]["label"], "客生主")
self.assertEqual(final_qi["six_qi"]["alignment"], "在泉同位")
self.assertEqual(controlled["framework"]["relations"]["guest_host"]["order"], "客胜为从")
def test_tianfu_and_suihui_use_traditional_year_positions(self):
taiyi = build_five_phase_field("2038-06-15")
non_suihui = build_five_phase_field("2022-06-15")
self.assertEqual(
taiyi["framework"]["relations"]["annual_pattern"]["primary"],
"太乙天符",
)
self.assertFalse(
non_suihui["framework"]["relations"]["annual_pattern"]["is_suihui"]
)
def test_public_field_has_no_observation_hour(self):
field = build_five_phase_field("2026-07-18")
self.assertNotIn("time", field["pillars"])
self.assertNotIn("observation_time", field)
def test_sector_catalog_lists_all_rules_and_applies_manual_overrides(self):
field = build_five_phase_field(
"2026-07-18",
{"电力": "", "低空经济": ""},
)
groups = {item["element"]: item["industries"] for item in field["sector_catalog"]}
names = {
element: {item["name"]: item["classification_source"] for item in items}
for element, items in groups.items()
}
self.assertEqual(set(groups), {"", "", "", "", ""})
self.assertNotIn("电力", names[""])
self.assertEqual(names[""]["电力"], "manual")
self.assertEqual(names[""]["低空经济"], "manual")
self.assertEqual(sum(len(items) for items in groups.values()), 144)
if __name__ == "__main__":
unittest.main()
+133
View File
@@ -0,0 +1,133 @@
from __future__ import annotations
import json
import re
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
STATIC = ROOT / "static"
class FrontendBoundaryTests(unittest.TestCase):
def test_shared_api_is_the_only_application_fetch_exit(self) -> None:
fetch_files = []
for path in STATIC.rglob("*.js"):
if "vendor" in path.parts:
continue
if re.search(r"\bfetch\s*\(", path.read_text(encoding="utf-8")):
fetch_files.append(path.relative_to(STATIC).as_posix())
self.assertEqual(fetch_files, ["shared/api.js"])
def test_shared_dependencies_load_before_application(self) -> None:
html = (STATIC / "index.html").read_text(encoding="utf-8")
ui_position = html.index('/ui-core.js')
components_position = html.index('/shared/components.js')
pages_position = html.index('/pages.config.js')
runtime_position = html.index('/pages/runtime.js')
state_position = html.index('/shared/state.js')
api_position = html.index('/shared/api.js')
shell_position = html.index('/shared/shell.js')
app_position = html.index('/app.js')
self.assertLess(ui_position, components_position)
self.assertLess(components_position, pages_position)
self.assertLess(pages_position, state_position)
self.assertLess(pages_position, runtime_position)
self.assertLess(runtime_position, state_position)
self.assertLess(state_position, api_position)
self.assertLess(api_position, shell_position)
self.assertLess(shell_position, app_position)
def test_application_state_is_created_through_shared_boundary(self) -> None:
app = (STATIC / "app.js").read_text(encoding="utf-8")
self.assertIn("const state = window.XiaobaiState.create({", app)
self.assertNotIn("const state = {", app)
def test_runtime_page_registry_matches_governance_registry(self) -> None:
expected = json.loads(
(ROOT / "config" / "pages.config.json").read_text(encoding="utf-8")
)["pages"]
runtime = (STATIC / "pages.config.js").read_text(encoding="utf-8")
rows = re.findall(
r'^\s*\["([^"]+)", "([^"]+)", "([^"]+)", "([^"]+)", "([^"]+)", (true|false)\],$',
runtime,
re.MULTILINE,
)
actual = [
{
"id": row[0],
"title": row[1],
"feature": row[2],
"group": row[3],
"access": row[4],
"default": row[5] == "true",
"desktop_scroll": "page",
"mobile_layout": "dedicated",
}
for row in rows
]
self.assertEqual(actual, expected)
def test_shell_owns_navigation_and_page_mounting(self) -> None:
app = (STATIC / "app.js").read_text(encoding="utf-8")
shell = (STATIC / "shared" / "shell.js").read_text(encoding="utf-8")
self.assertNotIn("function syncNavigationState", app)
self.assertNotIn("function initializeApplicationShell", app)
self.assertIn("function syncNavigation(viewId)", shell)
self.assertIn("function mount(viewId, mountOptions = {})", shell)
self.assertIn("function openModalDialog(dialog)", shell)
self.assertNotIn('document.querySelectorAll(".module-tab").forEach', app)
def test_every_registered_view_has_one_feature_page_module(self) -> None:
html = (STATIC / "index.html").read_text(encoding="utf-8")
runtime_position = html.index('/pages/runtime.js')
app_position = html.index('/app.js')
expected = {
page["id"]: page["feature"]
for page in json.loads(
(ROOT / "config" / "pages.config.json").read_text(encoding="utf-8")
)["pages"]
}
expected["screenerTrackingView"] = "screener"
actual: dict[str, str] = {}
for path in (STATIC / "pages").glob("*/page.js"):
script_url = f'/pages/{path.parent.name}/page.js'
self.assertIn(script_url, html)
self.assertLess(runtime_position, html.index(script_url))
self.assertLess(html.index(script_url), app_position)
script = path.read_text(encoding="utf-8")
for match in re.finditer(
r'XiaobaiPageModules\.register\("([^"]+)",\s*\[(.*?)\]',
script,
re.DOTALL,
):
feature = match.group(1)
for view_id in re.findall(r'"([A-Za-z][A-Za-z0-9]+)"', match.group(2)):
self.assertNotIn(view_id, actual)
actual[view_id] = feature
self.assertEqual(actual, expected)
def test_page_lifecycle_is_owned_outside_application_monolith(self) -> None:
app = (STATIC / "app.js").read_text(encoding="utf-8")
runtime = (STATIC / "pages" / "runtime.js").read_text(encoding="utf-8")
start = app.index("function openView(")
end = app.index("\nfunction initializeAutoTableSorting", start)
open_view = app[start:end]
self.assertIn("pageModules.beforeMount(viewId, previousView);", open_view)
self.assertIn("pageModules.afterMount(viewId, previousView);", open_view)
self.assertNotRegex(open_view, r'viewId\s*[!=]==?\s*"')
self.assertIn("function beforeMount(viewId, previousView)", runtime)
self.assertIn("function afterMount(viewId, previousView)", runtime)
def test_shared_empty_state_component_is_used_by_multiple_features(self) -> None:
components = (STATIC / "shared" / "components.js").read_text(encoding="utf-8")
app = (STATIC / "app.js").read_text(encoding="utf-8")
self.assertIn("function emptyStateHtml(message, options = {})", components)
self.assertIn("function renderEmptyState(target, message, options = {})", components)
self.assertGreaterEqual(app.count("renderEmptyState("), 8)
self.assertGreaterEqual(app.count("emptyStateHtml("), 8)
if __name__ == "__main__":
unittest.main()
+331
View File
@@ -0,0 +1,331 @@
from __future__ import annotations
import re
import unittest
from html.parser import HTMLParser
from pathlib import Path
STATIC_DIR = Path(__file__).resolve().parents[1] / "static"
class IdCollector(HTMLParser):
def __init__(self) -> None:
super().__init__()
self.ids: list[str] = []
def handle_starttag(self, tag, attrs):
self.ids.extend(value for key, value in attrs if key == "id" and value)
class FrontendContractTests(unittest.TestCase):
@classmethod
def setUpClass(cls) -> None:
cls.html = (STATIC_DIR / "index.html").read_text(encoding="utf-8")
cls.script = (STATIC_DIR / "app.js").read_text(encoding="utf-8")
cls.shell = (STATIC_DIR / "shared" / "shell.js").read_text(encoding="utf-8")
cls.ui_core = (STATIC_DIR / "ui-core.js").read_text(encoding="utf-8")
cls.design_system = (STATIC_DIR / "design-system.css").read_text(encoding="utf-8")
cls.theme = (STATIC_DIR / "theme.css").read_text(encoding="utf-8")
cls.tokens = (STATIC_DIR / "shared" / "tokens.css").read_text(encoding="utf-8")
collector = IdCollector()
collector.feed(cls.html)
cls.ids = collector.ids
def test_html_ids_are_unique(self):
duplicates = sorted({item for item in self.ids if self.ids.count(item) > 1})
self.assertEqual(duplicates, [])
def test_literal_id_selectors_exist_in_html(self):
selectors = set(re.findall(r'querySelector\("#([A-Za-z][A-Za-z0-9_-]*)"\)', self.script))
selectors.update(re.findall(r'getElementById\("([A-Za-z][A-Za-z0-9_-]*)"\)', self.script))
selectors.update(re.findall(r'setText\("([A-Za-z][A-Za-z0-9_-]*)"', self.script))
missing = sorted(selectors - set(self.ids))
self.assertEqual(missing, [])
def test_all_primary_views_have_navigation_entries(self):
views = set(re.findall(r'id="([A-Za-z][A-Za-z0-9_-]*View|limitPool)" class="workspace-view', self.html))
internal_views = set(re.findall(
r'<section id="([A-Za-z][A-Za-z0-9_-]*View)" class="workspace-view[^"]*"[^>]*\bdata-internal-view\b',
self.html,
))
navigation = set(re.findall(r'data-view="([A-Za-z][A-Za-z0-9_-]*)"', self.html))
self.assertEqual(views - internal_views, navigation)
self.assertEqual(len(views - internal_views), 16)
self.assertEqual(internal_views, {"screenerTrackingView"})
def test_market_discovery_views_are_wired_end_to_end(self):
for view_id in ("auctionView", "themeLibraryView", "popularityView"):
self.assertIn(f'id="{view_id}"', self.html)
self.assertIn(f'data-view="{view_id}"', self.html)
for endpoint in ("/api/auction?", "/api/themes?", "/api/themes/detail?", "/api/popularity?"):
self.assertIn(endpoint, self.script)
for field in (
"auction_change", "auction_amount_million",
"auction_turnover_rate", "auction_volume_ratio",
):
self.assertIn(field, (STATIC_DIR.parent / "screener.py").read_text(encoding="utf-8"))
def test_wencai_workspace_is_not_exposed_and_mentor_hides_internal_quality_score(self):
self.assertNotIn('id="wencaiView"', self.html)
self.assertNotIn('data-view="wencaiView"', self.html)
for endpoint in ("/api/wencai", "/api/wencai/query", "/api/wencai/saved"):
self.assertNotIn(endpoint, self.script)
self.assertNotIn("${score}/${total}", self.script)
def test_auction_navigation_and_frontend_pools_follow_product_order(self):
rotation = self.html.index('data-view="rotationView"')
auction = self.html.index('data-view="auctionView"')
themes = self.html.index('data-view="themeLibraryView"')
self.assertLess(rotation, auction)
self.assertLess(auction, themes)
for dataset in ("focus", "watchlist", "all", "onePrice"):
self.assertIn(f'data-auction-dataset="{dataset}"', self.html)
dataset_positions = [self.html.index(f'data-auction-dataset="{dataset}"') for dataset in ("focus", "watchlist", "all", "onePrice")]
self.assertEqual(dataset_positions, sorted(dataset_positions))
for filter_name in ("all", "above", "matched", "below"):
self.assertIn(f'data-auction-filter="{filter_name}"', self.html)
self.assertNotIn('data-auction-filter="strong"', self.html)
self.assertNotIn('data-auction-filter="limit"', self.html)
self.assertIn('id="auctionThemeCarry"', self.html)
self.assertIn('id="auctionAmountTrend"', self.html)
self.assertNotIn('id="auctionNewsTitle"', self.html)
self.assertIn('id="auctionWorkspaceTitle"', self.html)
self.assertIn('id="auctionExpectationFilterbar"', self.html)
self.assertIn('id="auctionExpectationControls"', self.html)
self.assertNotIn('id="auctionAboveCount"', self.html)
self.assertNotIn('id="auctionMatchedCount"', self.html)
self.assertNotIn('id="auctionBelowCount"', self.html)
self.assertNotIn('class="auction-news-entry"', self.html)
def test_visual_renovation_keeps_required_product_controls(self):
for order in ("oldest", "latest"):
self.assertIn(f'data-rotation-order="{order}"', self.html)
self.assertIn('id="dragonProfilesButton"', self.html)
self.assertIn('id="sentimentHistoryBody"', self.html)
self.assertIn('id="sentimentPreviousPositive"', self.html)
self.assertIn('id="accountDropdown"', self.html)
self.assertIn('id="settingsButton"', self.html)
def test_screener_uses_progressive_strategy_editor(self):
for step in ("regime", "strategy", "run", "result"):
self.assertIn(f'data-screener-step="{step}"', self.html)
def test_screener_exposes_curated_and_quant_workspaces(self):
for mode in ("smart", "curated", "quant"):
self.assertIn(f'data-screener-mode="{mode}"', self.html)
self.assertIn(f'data-screener-panel="{mode}"', self.html)
for element_id in (
"curatedStrategyList", "quantFilterRows",
"quantScoreRows", "quantRunButton", "quantSaveButton",
):
self.assertIn(f'id="{element_id}"', self.html)
for removed_id in (
"curatedRunButton", "factorSyncButton", "screenerRunButton",
"changeStrategyButton",
):
self.assertNotIn(f'id="{removed_id}"', self.html)
self.assertIn("盘后自动候选池", self.html)
self.assertIn("自定义选股", self.html)
self.assertIn('id="strategyDrawer" class="strategy-drawer"', self.html)
self.assertIn('id="openStrategyDrawerButton"', self.html)
self.assertIn('id="closeStrategyDrawerButton"', self.html)
self.assertIn('id="activeStrategyDescription"', self.html)
self.assertIn('openStrategyDrawer("editor")', self.script)
for element_id in ("curatedSuitableEnvironment", "curatedFailureRisk"):
self.assertIn(f'id="{element_id}"', self.html)
self.assertIn("meta.suitable_environment", self.script)
self.assertIn("meta.failure_risk", self.script)
self.assertIn('mode === "curated" ? "暂无符合条件个股"', self.script)
def test_curated_library_explains_empty_signals_and_supports_school_views(self):
for element_id in ("curatedSchoolFilters", "curatedStrategyList"):
self.assertIn(f'id="{element_id}"', self.html)
for view in ("list", "grid"):
self.assertIn(f'data-curated-view="{view}"', self.html)
for school in ("基本面", "趋势", "短线", "动量"):
self.assertIn(school, self.script)
self.assertIn("curatedStrategyRunState", self.script)
self.assertIn("必需数据已完整,本日没有股票同时满足", self.script)
def test_dialogs_and_dark_table_hover_have_shared_safety_constraints(self):
redesign = (STATIC_DIR / "redesign-v2.css").read_text(encoding="utf-8")
self.assertIn(".settings-dialog:not(.heaven-reading-dialog)[open] { margin: auto; }", redesign)
self.assertIn("max-height: min(760px, calc(100dvh - 28px));", redesign)
self.assertIn('#reviewWorkspaceView .data-table tbody tr:hover td', self.theme)
self.assertIn('#reviewWorkspaceView .data-table tbody td', self.theme)
self.assertIn('#screenerView .screener-result-frame tbody tr:hover td:last-child', self.theme)
def test_global_toast_has_one_owner_and_cannot_stretch_between_insets(self):
styles = (STATIC_DIR / "styles.css").read_text(encoding="utf-8")
wentian = (STATIC_DIR / "wentian-v2.css").read_text(encoding="utf-8")
self.assertIn("#toast.toast {", styles)
self.assertIn("top: auto;", styles)
self.assertIn("left: auto;", styles)
self.assertIn("height: auto;", styles)
self.assertIn("#toast.toast[hidden] { display: none; }", styles)
self.assertNotIn(".toast{position:fixed", self.design_system)
self.assertNotRegex(wentian, r"(?m)^\.toast\s*\{")
def test_public_knowledge_editors_are_hidden_for_non_admins(self):
self.assertIn('document.querySelector("#reasonForm").hidden = !isAdmin;', self.script)
self.assertIn('document.querySelector("#sectorPhaseManager").hidden = !isAdmin;', self.script)
self.assertIn('const canManage = state.user?.role === "admin";', self.script)
def test_shared_ui_core_loads_before_application(self):
self.assertLess(
self.html.index('<script src="/ui-core.js"'),
self.html.index('<script src="/app.js'),
)
for function_name in (
"number", "clamp", "escapeHtml", "formatNumber", "formatTimestamp",
"displayCompactDate", "todayString", "localDateString", "parseLocalDate",
):
self.assertIn(f"function {function_name}(", self.ui_core)
def test_business_views_do_not_expose_engineering_source_labels(self):
prohibited = (
"Tushare 实时行情",
"Tushare 日K",
"SQLite 缓存",
"演示日K",
"rt_k 实时截面",
)
for label in prohibited:
self.assertNotIn(label, self.script)
def test_stock_hover_preview_always_uses_latest_market_context(self):
start = self.script.index("async function showStockPreview")
end = self.script.index("async function showEntityPreview", start)
preview_loader = self.script[start:end]
self.assertIn('const cacheKey = `${code}:latest`;', preview_loader)
self.assertIn('/preview`', preview_loader)
self.assertNotIn("trade_date", preview_loader)
self.assertNotIn("elements.tradeDate.value", preview_loader)
def test_daily_rising_candles_are_fully_hollow_without_crossing_wicks(self):
start = self.script.index("function drawCandlestick")
end = self.script.index("function drawPriceChart", start)
candle = self.script[start:end]
self.assertIn("context.lineTo(x, bodyTop);", candle)
self.assertIn("context.moveTo(x, bodyBottom);", candle)
self.assertIn("context.lineTo(x, lowY);", candle)
self.assertIn("context.fillStyle = palette.background;", candle)
self.assertIn("context.strokeRect(bodyLeft, bodyTop, candleWidth, bodyHeight);", candle)
self.assertNotIn("context.lineTo(x, lowY);\n context.stroke();\n const openY", candle)
def test_stock_hover_intraday_draws_average_without_source_label(self):
start = self.script.index("function drawIntradayCanvas")
end = self.script.index("function drawDailyPreviewChart", start)
chart = self.script[start:end]
self.assertIn("point.average", chart)
self.assertIn("context.strokeStyle = palette.average;", chart)
self.assertIn('intraday_trade_date || payload.meta?.trade_date', self.script)
self.assertIn('(payload.intraday || []).length ? "最新分时 · 1分钟"', self.script)
def test_hover_prefers_daily_and_intraday_uses_centered_zero_axis(self):
self.assertIn('stockPreviewChart: "daily"', self.script)
self.assertIn('state.stockPreviewChart = "daily";', self.script)
self.assertIn('selectStockPreviewChart("daily");', self.script)
start = self.script.index("function drawIntradayCanvas")
end = self.script.index("function drawIntradayPreviewChart", start)
intraday = self.script[start:end]
self.assertIn("Math.abs(maximum - previousClose)", intraday)
self.assertIn("Math.abs(previousClose - minimum)", intraday)
self.assertIn('context.fillText("0.00%"', intraday)
self.assertIn('label: "09:30"', intraday)
self.assertIn('label: "11:30 / 13:00"', intraday)
self.assertIn('label: "15:00"', intraday)
self.assertIn("intradayMinuteOffset(points[index]?.time) / 240", intraday)
def test_detail_dialogs_offer_lazy_daily_and_intraday_modes(self):
self.assertIn('data-stock-detail-chart="daily"', self.html)
self.assertIn('data-stock-detail-chart="intraday"', self.html)
self.assertIn('data-entity-detail-chart="daily"', self.html)
self.assertIn('data-entity-detail-chart="intraday"', self.html)
self.assertIn('/api/chart/intraday?', self.script)
self.assertIn('drawIntradayCanvas(elements.priceChart', self.script)
self.assertIn('drawIntradayCanvas(elements.entityDetailChart', self.script)
self.assertIn('state.stockDetailChartMode === "intraday"', self.script)
self.assertIn('state.entityDetailChartMode === "intraday"', self.script)
def test_entity_daily_chart_uses_runtime_theme_palette(self):
start = self.script.index("function drawEntityDetailChart")
end = self.script.index("function clearEntityDetailChart", start)
chart = self.script[start:end]
self.assertIn("const palette = currentChartPalette();", chart)
self.assertIn("context.fillStyle = palette.background;", chart)
self.assertIn("context.fillStyle = palette.axis;", chart)
self.assertNotIn('context.fillStyle = "#6c7983";', chart)
def test_dark_mentor_tokens_and_sentiment_bottom_clearance_are_defined(self):
self.assertIn(":root[data-theme=\"dark\"] #mentorView {", self.theme)
self.assertIn("--mentor-ink: var(--text-primary);", self.theme)
self.assertIn("--mentor-sub: var(--text-secondary);", self.theme)
self.assertIn(":root[data-theme=\"dark\"] #mentorView .mentor-message {", self.theme)
self.assertIn("border-color: var(--line-soft);", self.theme)
self.assertIn("background: var(--surface-subtle);", self.theme)
self.assertIn("box-shadow: none;", self.theme)
self.assertIn("#sentimentCycleView .sentiment-history-frame {", self.theme)
self.assertIn("margin-bottom: var(--card-gap);", self.theme)
self.assertIn("padding-bottom: var(--card-gap);", self.theme)
self.assertIn("--sentiment-history-max-height: 510px;", self.tokens)
self.assertIn("max-height:var(--sentiment-history-max-height);", self.design_system)
self.assertIn("overflow:auto;", self.design_system)
def test_theme_switch_is_atomic_and_theme_library_loading_surface_is_dark_safe(self):
self.assertIn('typeof document.startViewTransition === "function"', self.script)
self.assertIn('root.classList.add("theme-switching")', self.script)
self.assertIn('root.classList.remove("theme-switching")', self.script)
self.assertIn("clearThemeTransitionEffects();", self.script)
self.assertIn("redrawThemeSensitiveVisuals();", self.script)
self.assertIn(":root.theme-switching *", self.theme)
self.assertIn("::view-transition-old(root)", self.theme)
self.assertIn(".theme-detail-empty-v2,", self.theme)
def test_membership_copy_includes_review_assistant_access(self):
self.assertIn("复盘助手仅对会员开放", self.html)
self.assertIn("智能选股、问师、问天、复盘助手等智能功能", self.html)
self.assertIn("自选股、复盘记录与交易日志", self.html)
self.assertIn("每日智能分析额度", self.html)
def test_review_assistant_uses_the_same_member_gate_pattern(self):
self.assertIn('id="assistantMemberGate" class="member-gate assistant-member-gate"', self.html)
self.assertIn('id="assistantMemberContent" class="assistant-member-content"', self.html)
self.assertIn('elements.assistantDialog.classList.toggle("member-locked", !unlocked);', self.script)
self.assertIn('button.disabled = !unlocked || state.assistantLoading;', self.script)
def test_trade_log_editor_is_dialog_based(self):
self.assertIn('id="openTradeLogDialog"', self.html)
self.assertIn('id="tradeLogDialog" class="settings-dialog trade-log-dialog"', self.html)
self.assertIn('openModalDialog(elements.tradeLogDialog)', self.script)
self.assertIn('document.querySelectorAll("dialog[open]")', self.shell)
self.assertIn('renderTradeLog();\n closeTradeLogDialog();', self.script)
def test_review_workspace_exposes_complete_watchlist_and_three_part_journal(self):
for label in (
"今日涨幅", "5日涨幅", "竞价关注(分)", "跟踪备注", "添加自选",
"今日盘面一句话", "今日做对了什么 / 做错了什么", "明日策略",
):
self.assertIn(label, self.html)
for element_id in (
"watchlistDialog", "watchlistSearchInput", "watchlistRemark",
"journalSummary", "journalContent", "journalPlan",
):
self.assertIn(f'id="{element_id}"', self.html)
self.assertIn('summary: document.querySelector("#journalSummary").value', self.script)
self.assertIn('return_5d', self.script)
self.assertIn('attention_score', self.script)
def test_heaven_interpretations_use_one_dialog_and_history_tabs(self):
self.assertIn('id="heavenReadingDialog"', self.html)
self.assertIn('data-heaven-reading-tab="current"', self.html)
self.assertIn('data-heaven-reading-tab="history"', self.html)
for button_id in ("historyTrendButton", "historyFortuneButton", "historyHeartButton"):
self.assertIn(f'id="{button_id}"', self.html)
self.assertIn('state.heavenInterpretations.fortune = payload.daily_fortune_reading || "";', self.script)
self.assertIn('if (existing) {\n openHeavenReading(mode, { loading: false });', self.script)
if __name__ == "__main__":
unittest.main()
+98
View File
@@ -0,0 +1,98 @@
from __future__ import annotations
import unittest
from pathlib import Path
from server import DashboardService
class SearchDatabaseStub:
def __init__(self) -> None:
self.directory = {
"schema_version": 2,
"items": [
{
"id": "881107.TI",
"code": "881107.TI",
"name": "油气开采及服务",
"type": "sector",
"subtitle": "行业板块",
"member_count": 19,
},
{
"id": "885728.TI",
"code": "885728.TI",
"name": "人工智能",
"type": "theme",
"subtitle": "概念题材",
"member_count": 1079,
},
],
}
def get_data_snapshot(self, kind: str, cache_key: str):
if (kind, cache_key) == ("search_directory", "ths"):
return self.directory
return None
@staticmethod
def search_stock_master(query: str, limit: int = 12):
if query in {"002141", "贤丰控股"}:
return [
{
"ts_code": "002141.SZ",
"code": "002141",
"name": "贤丰控股",
"industry": "元件",
"market": "主板",
"list_date": "20071228",
}
]
return []
class GlobalSearchTests(unittest.TestCase):
def setUp(self) -> None:
self.service = DashboardService.__new__(DashboardService)
self.service.database = SearchDatabaseStub()
self.service._system_credentials = {"tushare_token": ""}
def test_search_groups_stock_sector_theme_and_index(self):
stock = self.service.search_entities("002141", "2026-07-22")
sector = self.service.search_entities("油气", "2026-07-22")
theme = self.service.search_entities("人工智能", "2026-07-22")
index = self.service.search_entities("上证指数", "2026-07-22")
self.assertEqual(stock["groups"]["stocks"][0]["name"], "贤丰控股")
self.assertEqual(stock["groups"]["stocks"][0]["industry"], "元件")
self.assertEqual(sector["groups"]["sectors"][0]["code"], "881107.TI")
self.assertEqual(theme["groups"]["themes"][0]["code"], "885728.TI")
self.assertEqual(index["groups"]["indices"][0]["code"], "000001.SH")
def test_empty_query_returns_all_groups_without_remote_lookup(self):
result = self.service.search_entities("", "2026-07-22")
self.assertEqual(
result["groups"],
{"stocks": [], "sectors": [], "themes": [], "indices": []},
)
def test_frontend_reuses_full_stock_detail_and_renders_market_daily_k(self):
static_dir = Path(__file__).resolve().parents[1] / "static"
html = (static_dir / "index.html").read_text(encoding="utf-8")
script = (static_dir / "app.js").read_text(encoding="utf-8")
self.assertIn('id="globalSearchButton"', html)
self.assertIn('id="globalSearchDialog"', html)
self.assertIn('id="entityDetailDialog"', html)
self.assertIn("行情走势", html)
self.assertIn('data-entity-detail-chart="daily"', html)
self.assertIn('data-entity-detail-chart="intraday"', html)
self.assertIn('event.key.toLowerCase() !== "k"', script)
self.assertIn('openStock(item.id, { code: item.code', script)
self.assertNotIn('include_notes', script)
self.assertIn('const candles = (series || [])', script)
self.assertIn('renderStockNotes(payload.notes || [])', script)
if __name__ == "__main__":
unittest.main()
+110
View File
@@ -0,0 +1,110 @@
from __future__ import annotations
import json
import unittest
from pathlib import Path
from tools.build_api_registry import build as build_api_registry
from tools.build_architecture_inventory import build as build_architecture_inventory
ROOT = Path(__file__).resolve().parents[1]
CONFIG = ROOT / "config"
def load(name: str) -> dict:
return json.loads((CONFIG / name).read_text(encoding="utf-8"))
class GovernanceRegistryTests(unittest.TestCase):
def setUp(self) -> None:
self.features = load("features.config.json")
self.pages = load("pages.config.json")
self.api = load("api.config.json")
self.data = load("data-fields.config.json")
self.quality = load("data-quality.config.json")
self.jobs = load("jobs.config.json")
def test_features_are_unique_and_use_declared_roles(self) -> None:
roles = set(self.features["roles"])
items = self.features["features"]
ids = [item["id"] for item in items]
self.assertEqual(len(ids), len(set(ids)))
self.assertTrue(all(item["access"] in roles for item in items))
self.assertTrue(all(isinstance(item["enabled"], bool) for item in items))
def test_page_registry_matches_the_current_primary_navigation(self) -> None:
inventory_pages = build_architecture_inventory()["pages"]
registered = self.pages["pages"]
self.assertEqual(
[(item["id"], item["title"]) for item in registered],
[(item["id"], item["title"]) for item in inventory_pages],
)
feature_ids = {item["id"] for item in self.features["features"]}
self.assertTrue(all(page["feature"] in feature_ids for page in registered))
self.assertEqual(sum(bool(page.get("default")) for page in registered), 1)
def test_api_registry_is_current_and_owned(self) -> None:
self.assertEqual(self.api, build_api_registry())
feature_ids = {item["id"] for item in self.features["features"]}
roles = set(self.features["roles"])
keys = []
for route in self.api["routes"]:
keys.append((route["method"], route["path"], route["match"]))
self.assertIn(route["feature"], feature_ids)
self.assertIn(route["access"], roles)
self.assertEqual(len(keys), len(set(keys)))
def test_public_routes_are_explicitly_limited(self) -> None:
public = {
(item["method"], item["path"])
for item in self.api["routes"]
if item["access"] == "public"
}
self.assertEqual(
public,
{
("GET", "/api/health"),
("POST", "/api/auth/login"),
("POST", "/api/auth/register"),
},
)
def test_calculation_datasets_use_approved_providers(self) -> None:
providers = self.data["providers"]
ids = []
for dataset in self.data["datasets"]:
ids.append(dataset["id"])
self.assertIn(dataset["primary"], providers)
for fallback in dataset.get("fallbacks", []):
self.assertIn(fallback, providers)
if dataset["usage"] == "calculation":
self.assertTrue(providers[dataset["primary"]]["calculation_allowed"])
if dataset["primary"] == "unresolved":
self.assertEqual(dataset["usage"], "blocked")
self.assertEqual(len(ids), len(set(ids)))
def test_every_dataset_has_one_quality_rule_and_known_unit_profile(self) -> None:
dataset_ids = {item["id"] for item in self.data["datasets"]}
rules = self.quality["datasets"]
self.assertEqual(set(rules), dataset_ids)
profiles = set(self.quality["unit_profiles"])
self.assertTrue(
all(str(rule.get("unit_profile") or "none") in profiles for rule in rules.values())
)
def test_background_jobs_have_complete_unique_runtime_contracts(self) -> None:
jobs = self.jobs["jobs"]
ids = [item["id"] for item in jobs]
self.assertEqual(len(ids), len(set(ids)))
for item in jobs:
self.assertTrue(item["schedule"])
self.assertTrue(item["input_date_policy"])
self.assertTrue(item["lock_key"])
self.assertGreater(int(item["timeout_seconds"]), 0)
self.assertGreater(int(item["max_attempts"]), 0)
self.assertTrue(item["output_version"])
if __name__ == "__main__":
unittest.main()
+129
View File
@@ -0,0 +1,129 @@
from __future__ import annotations
import tempfile
import threading
import unittest
from pathlib import Path
from unittest.mock import patch
from database import ReviewDatabase
from server import DashboardService
class HeavenReadingTests(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("heaven_owner", "salt", "hash")
self.other = self.database.create_user("heaven_other", "salt", "hash")
def tearDown(self) -> None:
self.temp.cleanup()
def save(self, user_id: int, dedupe_key: str = "fortune:20260723") -> dict:
return self.database.save_heaven_reading(
user_id,
"fortune",
"20260723",
"2026-07-23 观气",
"丙午年 · 乙未月 · 己丑日",
"当日解运结果",
{"calendar_date": "20260723"},
dedupe_key,
)
def test_history_is_private_and_delete_requires_ownership(self):
reading = self.save(self.owner["id"])
self.assertEqual(
self.database.list_heaven_readings(self.owner["id"], "fortune")[0]["answer"],
"当日解运结果",
)
self.assertEqual(
self.database.list_heaven_readings(self.other["id"], "fortune"), []
)
self.assertFalse(
self.database.delete_heaven_reading(self.other["id"], reading["id"])
)
self.assertTrue(
self.database.delete_heaven_reading(self.owner["id"], reading["id"])
)
def test_fortune_dedupe_keeps_the_first_successful_result(self):
first = self.save(self.owner["id"])
second = self.database.save_heaven_reading(
self.owner["id"],
"fortune",
"20260723",
"replacement",
"replacement",
"不应覆盖",
{},
"fortune:20260723",
)
self.assertEqual(second["id"], first["id"])
self.assertEqual(second["answer"], "当日解运结果")
def test_fortune_interpretation_reuses_saved_result_before_llm(self):
existing = self.save(self.owner["id"])
service = DashboardService.__new__(DashboardService)
service.database = self.database
service._request_context = threading.local()
service._request_context.user_id = self.owner["id"]
with patch.object(service, "heaven_setup") as setup, patch.object(
service, "_call_heaven_agent"
) as call_agent:
result = service.heaven_interpret(
{"mode": "fortune", "trade_date": "2026-07-23"}
)
setup.assert_not_called()
call_agent.assert_not_called()
self.assertTrue(result["reused"])
self.assertEqual(result["reading"]["id"], existing["id"])
self.assertEqual(result["answer"], "当日解运结果")
def test_legacy_truncated_fortune_is_regenerated_once(self):
self.database.save_heaven_reading(
self.owner["id"],
"fortune",
"20260723",
"2026-07-23 观气",
"丙午年 · 乙未月 · 己丑日",
"旧版解运结果被截断……",
{"calendar_date": "20260723"},
"fortune:20260723",
)
service = DashboardService.__new__(DashboardService)
service.database = self.database
service._request_context = threading.local()
service._request_context.user_id = self.owner["id"]
setup = {
"calendar_date": "20260723",
"field": {
"sector_catalog": [],
"balance": [{"element": ""}],
"pillars": {"year": "丙午", "month": "乙未", "day": "己丑"},
},
}
with patch.object(service, "heaven_setup", return_value=setup), patch.object(
service, "account_personal_field", return_value={}
), patch.object(
service,
"_call_heaven_agent",
return_value=({"answer": "完整的当日解运结果", "model": "test", "latency_ms": 1}, "primary"),
):
result = service.heaven_interpret(
{"mode": "fortune", "trade_date": "2026-07-23"}
)
self.assertFalse(result["reused"])
self.assertEqual(result["reading"]["answer"], "完整的当日解运结果")
self.assertEqual(
len(self.database.list_heaven_readings(self.owner["id"], "fortune")), 1
)
if __name__ == "__main__":
unittest.main()
+382
View File
@@ -0,0 +1,382 @@
from __future__ import annotations
import http.client
import json
import unittest
from unittest.mock import MagicMock, patch
from heaven_engine import _market_line_scores, build_manual_market_hexagram
from realtime_aggregator import WebRealtimeAggregator
from server import DashboardService
from tushare_client import (
TushareClient,
_filter_members_by_listing,
_sector_coverage_issue,
)
class HeavenMarketLineTests(unittest.TestCase):
def test_index_external_uses_only_current_index_change(self):
dashboard = {
"overview": {
"up_count": 1740,
"down_count": 3710,
"amount_billion": 27180.8,
"sentiment_score": 27,
"seal_rate": 57.3,
"limit_up_count": 55,
"limit_down_count": 267,
},
"sectors": [],
"sector_rotation": [],
}
index_context = {
"aggregate": {
"average_pct_chg": 0.187,
"average_return_5d": -5.606,
}
}
scores = _market_line_scores(
dashboard,
[{"amount_billion": 25000}],
index_context,
{},
{},
[],
)
self.assertAlmostEqual(scores[5]["score"], 0.187 / 3)
self.assertGreater(scores[5]["score"], 0)
self.assertIn("不参与外显阴阳", scores[5]["evidence"][1])
def test_manual_calibration_preserves_six_lines_and_marks_user_evidence(self):
chart = build_manual_market_hexagram(
[8, 7, 6, 9, 8, 7],
"20260722",
{"name": "元件", "taxonomy": "sw_l2"},
{"code": "002141", "name": "贤丰控股"},
{},
"人工核对",
)
self.assertTrue(chart["manual_calibration"])
self.assertEqual([line["value"] for line in chart["hexagram"]["lines"]], [8, 7, 6, 9, 8, 7])
self.assertEqual(chart["hexagram"]["moving_lines"], [3, 4])
self.assertIn("用户手动校准", chart["hexagram"]["lines"][0]["evidence"][0])
def test_quantitative_supplement_repairs_only_failed_lines_and_recalculates(self):
trade_date = "20260722"
dashboard = {
"overview": {
"sentiment_score": 32,
"seal_rate": 48,
"amount_billion": 16500,
"up_count": 1800,
"down_count": 3500,
"limit_up_count": 35,
"limit_down_count": 192,
},
"limits": [{"amount_billion": 12}, {"amount_billion": 25}],
"sectors": [],
"sector_rotation": [],
}
history = [{"amount_billion": 16000}, {"amount_billion": 15800}]
index_context = {
"trade_date": trade_date,
"source": "tushare",
"realtime": False,
"precise": False,
"indices": [
{"ts_code": code, "trade_date": "20260721", "pct_chg": 0.1}
for code in ("000001.SH", "399001.SZ", "399006.SZ")
],
}
sector = {"taxonomy": "sw_l2", "precise": False, "error": "行业日线尚未返回"}
stock = {
"code": "002141", "name": "贤丰控股", "trade_date": trade_date,
"data_source": "tushare", "realtime": False, "precise": True,
"amount_billion": 20, "turnover_rate": 8.5, "seal_amount_million": 0,
"open_times": 0, "change": 2.4, "streak": 0, "status": "普通",
}
automatic = DashboardService._heaven_line_checks(
trade_date, dashboard, history, index_context, sector, stock, "closed", {}
)
self.assertEqual(
[item["line"] for item in automatic if not item["passed"]], [3, 4, 6]
)
manual = {
"sector_name": "元件", "sector_up_count": 18, "sector_down_count": 42,
"sector_coverage": 96, "sector_member_equal_change": -2.2,
"sector_change": -2.6, "sector_leading_pct": 3.1,
"index_sh_change": -0.9, "index_sz_change": -1.4, "index_cy_change": -1.8,
}
merged = DashboardService._apply_heaven_manual_data(
dashboard, index_context, sector, stock, manual, "closed", trade_date, "002141"
)
repaired = DashboardService._heaven_line_checks(
trade_date, merged[0], history, merged[1], merged[2], merged[3], "closed", manual
)
self.assertTrue(all(item["passed"] for item in repaired))
self.assertEqual(
[item["line"] for item in repaired if item["status"] == "manual"], [3, 4, 6]
)
self.assertEqual(repaired[5]["line_value"], 8)
self.assertAlmostEqual(repaired[5]["score"], (-0.9 - 1.4 - 1.8) / 3 / 3, places=3)
def test_sector_inner_and_outer_have_independent_quality_gates(self):
trade_date = "20260722"
dashboard = {
"overview": {
"sentiment_score": 30, "seal_rate": 50, "amount_billion": 15000,
"up_count": 2000, "down_count": 3000,
"limit_up_count": 40, "limit_down_count": 80,
},
"limits": [], "sectors": [], "sector_rotation": [],
}
history = [{"amount_billion": 14800}, {"amount_billion": 14900}]
indices = {
"trade_date": trade_date, "source": "tushare", "realtime": False,
"precise": True,
"indices": [
{"ts_code": code, "trade_date": trade_date, "pct_chg": -1}
for code in ("000001.SH", "399001.SZ", "399006.SZ")
],
"aggregate": {"average_pct_chg": -1},
}
sector = {
"name": "元件", "code": "801083.SI", "taxonomy": "sw_l2",
"trade_date": trade_date, "realtime": False, "finalized": True,
"inner_precise": True, "outer_precise": False, "precise": False,
"coverage": 98, "up_count": 19, "down_count": 46,
"member_equal_change": -2.15, "leading_pct": 5.2,
"outer_error": "申万日线尚未发布", "source": "tushare_member_daily",
}
stock = {
"code": "002141", "trade_date": trade_date, "precise": True,
"realtime": False, "data_source": "tushare", "amount_billion": 10,
"turnover_rate": 5, "seal_amount_million": 0, "open_times": 0,
"change": 2, "streak": 0, "status": "普通",
}
checks = DashboardService._heaven_line_checks(
trade_date, dashboard, history, indices, sector, stock, "closed", {}
)
self.assertTrue(checks[2]["passed"])
self.assertFalse(checks[3]["passed"])
self.assertIn("申万日线尚未发布", checks[3]["reasons"])
manual = {"sector_change": -3.85}
merged = DashboardService._apply_heaven_manual_data(
dashboard, indices, sector, stock, manual, "closed", trade_date, "002141"
)
repaired = DashboardService._heaven_line_checks(
trade_date, merged[0], history, merged[1], merged[2], merged[3], "closed", manual
)
self.assertTrue(repaired[3]["passed"])
self.assertEqual(repaired[3]["status"], "manual")
self.assertEqual(
[field["key"] for field in repaired[3]["fields"] if field["manual"]],
["sector_change"],
)
missing_leader_sector = {**merged[2], "leading_pct": None}
still_blocked = DashboardService._heaven_line_checks(
trade_date,
merged[0],
history,
merged[1],
missing_leader_sector,
merged[3],
"closed",
manual,
)
self.assertFalse(still_blocked[3]["passed"])
self.assertIn("需补充:行业领涨股涨跌幅", still_blocked[3]["reasons"])
class ShenwanMembershipTests(unittest.TestCase):
def test_confirmed_delisted_members_are_removed_for_the_target_date(self):
members = [
{"ts_code": "601318.SH", "name": "中国平安"},
{"ts_code": "601319.SH", "name": "中国人保"},
{"ts_code": "000627.SZ", "name": "退市成员"},
{"ts_code": "999999.SZ", "name": "状态未知成员"},
]
reference = {
"601318.SH": {"list_date": "20070228", "delist_date": ""},
"601319.SH": {"list_date": "20181030", "delist_date": ""},
"000627.SZ": {"list_date": "19961112", "delist_date": "20250930"},
}
eligible, excluded = _filter_members_by_listing(
members, reference, "20260723"
)
self.assertEqual(
[item["ts_code"] for item in eligible],
["601318.SH", "601319.SH", "999999.SZ"],
)
self.assertEqual(excluded[0]["ts_code"], "000627.SZ")
self.assertEqual(excluded[0]["reason"], "目标日期前已退市")
def test_member_is_kept_for_dates_before_its_delisting(self):
members = [{"ts_code": "000627.SZ", "name": "历史有效成员"}]
reference = {
"000627.SZ": {"list_date": "19961112", "delist_date": "20250930"}
}
eligible, excluded = _filter_members_by_listing(
members, reference, "20250929"
)
self.assertEqual(eligible, members)
self.assertEqual(excluded, [])
def test_sector_coverage_gate_adapts_to_member_count(self):
self.assertEqual(_sector_coverage_issue(5, 5, 100), "")
self.assertIn("全部可解释", _sector_coverage_issue(5, 4, 80))
self.assertEqual(_sector_coverage_issue(5, 4, 100, 5), "")
self.assertEqual(_sector_coverage_issue(10, 9, 90), "")
self.assertIn("至少90%", _sector_coverage_issue(9, 8, 88.9))
self.assertIn("最多缺1只", _sector_coverage_issue(20, 18, 90))
self.assertEqual(_sector_coverage_issue(50, 45, 90), "")
self.assertIn("低于90%", _sector_coverage_issue(50, 44, 88))
@patch.object(TushareClient, "query")
def test_confirmed_suspension_explains_a_missing_quote(self, query: MagicMock):
TushareClient._suspension_cache.clear()
query.return_value = [{
"ts_code": "601319.SH",
"suspend_date": "20260720",
"resume_date": "20260725",
"suspend_reason": "重大事项",
}]
members = [
{"ts_code": "601318.SH", "name": "中国平安"},
{"ts_code": "601319.SH", "name": "中国人保"},
]
suspended = TushareClient("token")._confirmed_suspended_members(
members, {"601318.SH"}, "20260723"
)
self.assertEqual(len(suspended), 1)
self.assertEqual(suspended[0]["ts_code"], "601319.SH")
self.assertEqual(suspended[0]["reason"], "重大事项")
@patch.object(TushareClient, "query")
def test_latest_effective_membership_wins_over_stale_is_new_row(self, query: MagicMock):
stale_y = {
"l1_code": "801010.SI", "l1_name": "农林牧渔",
"l2_code": "801018.SI", "l2_name": "动物保健Ⅱ",
"l3_code": "850181.SI", "l3_name": "动物保健Ⅲ",
"ts_code": "002141.SZ", "in_date": "20240730", "out_date": None, "is_new": "Y",
}
current_y = {
"l1_code": "801080.SI", "l1_name": "电子",
"l2_code": "801083.SI", "l2_name": "元件",
"l3_code": "850822.SI", "l3_name": "印制电路板",
"ts_code": "002141.SZ", "in_date": "20260701", "out_date": None, "is_new": "Y",
}
closed_n = {**stale_y, "out_date": "20260630", "is_new": "N"}
query.side_effect = lambda _api, params, _fields: (
[stale_y, current_y] if params["is_new"] == "Y" else [closed_n]
)
industry = TushareClient("token").sw_stock_industry("002141.SZ", "20260722")
self.assertEqual(industry["l2_code"], "801083.SI")
self.assertEqual(industry["l2_name"], "元件")
class RealtimeAggregatorTests(unittest.TestCase):
def setUp(self):
WebRealtimeAggregator._response_cache.clear()
@staticmethod
def _response(payload: dict) -> MagicMock:
response = MagicMock()
response.headers.get.return_value = "application/json"
response.read.return_value = json.dumps(payload).encode("utf-8")
context = MagicMock()
context.__enter__.return_value = response
return context
@patch("realtime_aggregator.urllib.request.urlopen")
def test_transport_failure_is_retried(self, urlopen: MagicMock):
urlopen.side_effect = [
http.client.RemoteDisconnected("temporary disconnect"),
self._response({"rc": 0, "data": {"diff": []}}),
]
aggregator = WebRealtimeAggregator(retry_delay_seconds=0)
payload = aggregator._get_json("https://example.test", {}, "https://example.test")
self.assertEqual(payload["rc"], 0)
self.assertEqual(urlopen.call_count, 2)
@patch("realtime_aggregator.urllib.request.urlopen")
def test_recent_success_is_used_after_retries_fail(self, urlopen: MagicMock):
aggregator = WebRealtimeAggregator(retry_delay_seconds=0)
urlopen.return_value = self._response({"rc": 0, "data": {"diff": []}})
aggregator._get_json("https://example.test", {}, "https://example.test")
urlopen.side_effect = http.client.RemoteDisconnected("temporary disconnect")
payload = aggregator._get_json("https://example.test", {}, "https://example.test")
self.assertIn("_aggregate_cache", payload)
self.assertEqual(urlopen.call_count, 4)
@patch.object(WebRealtimeAggregator, "_get_text")
def test_tencent_indices_include_verifiable_quote_times(self, get_text: MagicMock):
def quote_line(
symbol: str,
name: str,
code: str,
price: str,
previous_close: str,
quote_time: str,
change_amount: str,
change: str,
amount: str,
) -> str:
fields = [""] * 38
fields[1] = name
fields[2] = code
fields[3] = price
fields[4] = previous_close
fields[5] = price
fields[30] = quote_time
fields[31] = change_amount
fields[32] = change
fields[33] = price
fields[34] = price
fields[37] = amount
return f'v_{symbol}="{"~".join(fields)}";'
get_text.return_value = (
'\n'.join(
[
quote_line("sh000001", "上证指数", "000001", "3796.28", "3764.15", "20260720155402", "32.13", "0.85", "129465190"),
quote_line("sz399001", "深证成指", "399001", "13610.23", "13706.88", "20260720155330", "-96.65", "-0.71", "140747525"),
quote_line("sz399006", "创业板指", "399006", "3443.10", "3428.63", "20260720155345", "14.47", "0.42", "67120487"),
]
),
0,
)
rows = WebRealtimeAggregator().tencent_indices()
self.assertEqual(len(rows), 3)
self.assertEqual(rows[0]["source"], "tencent_qt")
self.assertEqual(rows[0]["quote_time"][:10], "2026-07-20")
self.assertAlmostEqual(rows[0]["amount_billion"], 12946.52)
if __name__ == "__main__":
unittest.main()
+73
View File
@@ -0,0 +1,73 @@
from __future__ import annotations
import unittest
from tushare_client import TushareClient
class HotMoneyProfileClient(TushareClient):
def query(self, api_name, params=None, fields=""):
self.last_request = (api_name, params or {}, fields)
if api_name != "hm_list":
raise AssertionError(f"unexpected api: {api_name}")
return [
{
"name": "赵老哥",
"desc": "聚焦市场核心标的。",
"orgs": "华泰证券浙江分公司;银河证券绍兴",
},
{
"name": "炒股养家",
"desc": "",
"orgs": "华鑫证券上海宛平南路, 华鑫证券上海分公司",
},
{
"name": "赵老哥",
"desc": "重复记录不应覆盖首条档案。",
"orgs": "重复席位",
},
{"name": "", "desc": "无效记录", "orgs": ""},
]
class HotMoneyProfileTests(unittest.TestCase):
def test_directory_normalizes_profiles_and_organizations(self):
client = HotMoneyProfileClient("token")
payload = client.hot_money_profiles()
self.assertEqual(client.last_request[0], "hm_list")
self.assertEqual(client.last_request[2], "name,desc,orgs")
self.assertEqual(payload["meta"]["status"], "success")
self.assertEqual(payload["summary"], {
"profile_count": 2,
"described_count": 1,
"organization_count": 4,
})
self.assertEqual(
payload["profiles"][0]["organizations"],
["华泰证券浙江分公司", "银河证券绍兴"],
)
self.assertEqual(payload["profiles"][1]["organization_count"], 2)
self.assertEqual(
[item["id"] for item in payload["profiles"]],
["hot-money-profile-1", "hot-money-profile-2"],
)
def test_directory_parses_json_encoded_organization_lists(self):
client = TushareClient("token")
client.query = lambda *_args, **_kwargs: [
{
"name": "Profile",
"desc": "",
"orgs": '["Seat A", "Seat B", "Seat A"]',
}
]
payload = client.hot_money_profiles()
self.assertEqual(payload["profiles"][0]["organizations"], ["Seat A", "Seat B"])
self.assertEqual(payload["summary"]["organization_count"], 2)
if __name__ == "__main__":
unittest.main()
+48
View File
@@ -0,0 +1,48 @@
from __future__ import annotations
import unittest
from http import HTTPStatus
from api_access import ROUTES, required_role
from backend.http import correlation_id, normalize_error_payload
from tools.build_api_registry import build as build_api_registry
class HttpGovernanceTests(unittest.TestCase):
def test_runtime_registry_resolves_every_declared_route(self) -> None:
for item in build_api_registry()["routes"]:
path = (
item["path"]
.replace("(\\d{6})", "000001")
.replace("(\\d+)", "1")
.replace("(.+)", "sample")
)
resolved = ROUTES.resolve(item["method"], path)
self.assertIsNotNone(resolved, f"{item['method']} {path}")
self.assertEqual(resolved.feature, item["feature"])
self.assertEqual(resolved.access, item["access"])
def test_unknown_route_has_no_runtime_match(self) -> None:
self.assertIsNone(ROUTES.resolve("GET", "/api/not-registered"))
def test_compatibility_access_function_uses_runtime_registry(self) -> None:
self.assertEqual(required_role("GET", "/api/screener/setup"), "member")
self.assertEqual(required_role("POST", "/api/admin/settings"), "admin")
self.assertEqual(required_role("GET", "/api/dashboard"), "authenticated")
def test_errors_keep_legacy_field_and_add_stable_contract(self) -> None:
result = normalize_error_payload(
{"error": "invalid input"}, HTTPStatus.BAD_REQUEST, "request-123"
)
self.assertEqual(result["error"], "invalid input")
self.assertEqual(result["message"], "invalid input")
self.assertEqual(result["code"], "bad_request")
self.assertEqual(result["request_id"], "request-123")
def test_correlation_id_rejects_header_injection(self) -> None:
self.assertEqual(correlation_id("client-request-123"), "client-request-123")
self.assertRegex(correlation_id("bad\r\nheader"), r"^[0-9a-f]{32}$")
if __name__ == "__main__":
unittest.main()
+54
View File
@@ -0,0 +1,54 @@
from __future__ import annotations
import unittest
from ifind_client import IfindError, IfindHttpClient
class IfindClientTests(unittest.TestCase):
def test_table_rows_normalizes_single_table_payload(self):
rows = IfindHttpClient._table_rows(
{
"tables": {
"thscode": "300033.SZ",
"time": ["2026-07-28 09:30", "2026-07-28 09:31"],
"table": {"close": [10.1, 10.2], "amount": [100, 200]},
}
}
)
self.assertEqual(len(rows), 2)
self.assertEqual(rows[0]["thscode"], "300033.SZ")
self.assertEqual(rows[1]["time"], "2026-07-28 09:31")
self.assertEqual(rows[1]["close"], 10.2)
def test_table_rows_normalizes_wencai_list_payload(self):
rows = IfindHttpClient._table_rows(
{
"tables": [
{
"table": {
"股票代码": ["000001.SZ", "600000.SH"],
"股票简称": ["平安银行", "浦发银行"],
}
}
]
}
)
self.assertEqual([row["股票代码"] for row in rows], ["000001.SZ", "600000.SH"])
def test_display_date_rejects_invalid_values(self):
self.assertEqual(IfindHttpClient._display_date("20260728"), "2026-07-28")
with self.assertRaises(IfindError):
IfindHttpClient._display_date("2026-7-28")
def test_client_requires_credentials_before_request(self):
client = IfindHttpClient()
self.assertFalse(client.configured)
with self.assertRaises(IfindError):
client.real_time("000001.SH", ["latest"])
if __name__ == "__main__":
unittest.main()
+202
View File
@@ -0,0 +1,202 @@
from __future__ import annotations
import tempfile
import unittest
from datetime import date, datetime, timedelta, timezone
from pathlib import Path
from unittest.mock import patch
from chart_data_provider import EastmoneyChartClient, MarketChartClient
from database import ReviewDatabase
from market_insights import MarketInsightsService
from server import DashboardService
class FakeIfind:
configured = True
def history(self, codes, indicators, start_date, end_date, cache_ttl=0):
return [
{
"time": "2026-07-27",
"thscode": "000001.SZ",
"open": 10,
"high": 10.5,
"low": 9.8,
"close": 10.2,
"volume": 100,
"amount": 1_000_000,
},
{
"time": "2026-07-28",
"thscode": "000001.SZ",
"open": 10.2,
"high": 10.8,
"low": 10.1,
"close": 10.5,
"volume": 120,
"amount": 1_200_000,
},
]
def real_time(self, codes, indicators, cache_ttl=0):
return []
class FakeIfindStalePreopen(FakeIfind):
def history(self, codes, indicators, start_date, end_date, cache_ttl=0):
return [
*super().history(codes, indicators, start_date, end_date, cache_ttl),
{
"time": "2026-07-29",
"thscode": "000001.SZ",
"open": 10.5,
"high": 10.5,
"low": 10.5,
"close": 10.5,
"volume": 0,
"amount": 0,
},
]
def real_time(self, codes, indicators, cache_ttl=0):
return [
{
"time": "2026-07-28 15:00:00",
"open": 10.2,
"high": 10.8,
"low": 10.1,
"latest": 10.5,
"preClose": 10.2,
"volume": 120,
"amount": 1_200_000,
}
]
class FixedPreopenDatetime(datetime):
fixed_now = datetime(2026, 7, 29, 8, 45, tzinfo=timezone(timedelta(hours=8)))
@classmethod
def now(cls, tz=None):
return cls.fixed_now
class FakeIfindSnapshots:
configured = True
def __init__(self):
self.calls = []
def snapshots(self, codes, indicators, start_time, end_time, cache_ttl=0):
self.calls.append(
{
"codes": codes,
"indicators": indicators,
"start_time": start_time,
"end_time": end_time,
"cache_ttl": cache_ttl,
}
)
return [
{
"time": "2026-07-28 09:21:00",
"thscode": "000001.SZ",
"latest": 10.5,
"preClose": 10,
"volume": 2000,
"amount": 21000,
"bidSize1": 1200,
"askSize1": 800,
}
]
class FakeTushare:
pass
class IfindFeatureTests(unittest.TestCase):
def test_wencai_saved_queries_are_isolated_by_user(self):
with tempfile.TemporaryDirectory() as temporary:
database = ReviewDatabase(Path(temporary) / "review.db")
first = database.create_user("first-user", "salt", "hash")
second = database.create_user("second-user", "salt", "hash")
database.save_wencai_query(first["id"], "高质量", "ROE大于15%", "stock")
self.assertEqual(len(database.list_wencai_saved_queries(first["id"])), 1)
self.assertEqual(database.list_wencai_saved_queries(second["id"]), [])
def test_ifind_daily_chart_normalizes_change(self):
client = MarketChartClient(FakeIfind(), EastmoneyChartClient())
rows = client.stock_daily("000001", "20260728")
self.assertEqual(rows[-1]["trade_date"], "2026-07-28")
self.assertAlmostEqual(rows[-1]["change"], 2.9412, places=4)
def test_ifind_daily_chart_keeps_last_traded_bar_before_market_open(self):
client = MarketChartClient(FakeIfindStalePreopen(), EastmoneyChartClient())
with patch("chart_data_provider.datetime", FixedPreopenDatetime):
rows = client.stock_daily("000001", "20260729")
self.assertEqual(rows[-1]["trade_date"], "2026-07-28")
self.assertFalse(rows[-1].get("realtime", False))
def test_event_enrichment_keeps_blank_broken_reason_blank(self):
dashboard = {"broken": [{"code": "000001", "reason": "原原因"}]}
DashboardService._merge_ifind_event_enrichment(
dashboard,
{
"broken": {
"000001": {
"reason": "",
"first_time": "09:42:00",
"last_time": "",
"open_times": 3,
}
}
},
)
self.assertEqual(dashboard["broken"][0]["reason"], "原原因")
self.assertEqual(dashboard["broken"][0]["open_times"], 3)
def test_dynamic_auction_uses_ifind_snapshot_window_and_normalizes_rows(self):
with tempfile.TemporaryDirectory() as temporary:
database = ReviewDatabase(Path(temporary) / "review.db")
database.upsert_stock_master(
[
{
"ts_code": "000001.SZ",
"name": "Ping An Bank",
"industry": "Bank",
"market": "MainBoard",
"list_date": "19910403",
}
]
)
ifind = FakeIfindSnapshots()
service = MarketInsightsService(
database,
FakeTushare(),
now_provider=lambda: datetime(
2026, 7, 28, 9, 22, tzinfo=timezone(timedelta(hours=8))
),
ifind=ifind,
)
service._auction_candidates = lambda rows, baseline: (
[{"ts_code": "000001.SZ"}],
{},
[],
)
rows = service._dynamic_auction_rows("20260728", "20260727", 0)
self.assertEqual(ifind.calls[0]["start_time"], "2026-07-28 09:15:00")
self.assertEqual(ifind.calls[0]["end_time"], "2026-07-28 09:22:00")
self.assertEqual(rows[0]["ts_code"], "000001.SZ")
self.assertEqual(rows[0]["price"], 10.5)
self.assertEqual(rows[0]["snapshot_time"], "2026-07-28 09:21:00")
self.assertTrue(rows[0]["dynamic"])
if __name__ == "__main__":
unittest.main()
+73
View File
@@ -0,0 +1,73 @@
from __future__ import annotations
import tempfile
import threading
import unittest
from pathlib import Path
from backend.jobs import InProcessJobRunner, JobRegistry, SQLiteJobRunRepository
from database import ReviewDatabase
class JobRunnerTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.addCleanup(self.temporary.cleanup)
self.database = ReviewDatabase(Path(self.temporary.name) / "review.db")
self.repository = SQLiteJobRunRepository(self.database)
self.runner = InProcessJobRunner(JobRegistry.load(), self.repository)
def test_successful_idempotent_job_runs_once(self) -> None:
calls = []
self.assertTrue(
self.runner.run_inline(
"screener.automatic", "20260729:v8", lambda: calls.append("run")
)
)
self.assertFalse(
self.runner.run_inline(
"screener.automatic", "20260729:v8", lambda: calls.append("again")
)
)
self.assertEqual(calls, ["run"])
self.assertEqual(self.repository.recent(1)[0]["status"], "success")
def test_failure_is_persisted_and_can_be_retried_later(self) -> None:
def fail() -> None:
raise RuntimeError("provider unavailable")
self.assertTrue(self.runner.run_inline("market.refresh", "failed-key", fail))
failed = self.repository.recent(1)[0]
self.assertEqual(failed["status"], "failed")
self.assertEqual(failed["error_code"], "RuntimeError")
self.assertTrue(
self.runner.run_inline("market.refresh", "failed-key", lambda: None)
)
self.assertEqual(self.repository.recent(1)[0]["attempt"], 2)
def test_lock_key_rejects_concurrent_submission(self) -> None:
entered = threading.Event()
release = threading.Event()
def wait() -> None:
entered.set()
release.wait(2)
self.assertTrue(self.runner.submit("market.refresh", "first", wait))
self.assertTrue(entered.wait(1))
self.assertFalse(self.runner.submit("market.refresh", "second", lambda: None))
release.set()
self.assertTrue(self.runner.wait_for_idle())
def test_failed_status_payload_is_recorded_as_failure(self) -> None:
self.assertTrue(
self.runner.run_inline(
"screener.automatic", "reported-failure",
lambda: {"status": "failed", "error": "missing factors"},
)
)
self.assertEqual(self.repository.recent(1)[0]["status"], "failed")
if __name__ == "__main__":
unittest.main()
+118
View File
@@ -0,0 +1,118 @@
from __future__ import annotations
import unittest
from backend.llm import LLMGateway, LLMGatewayError
class ProviderFailure(RuntimeError):
pass
class FakeDatabase:
def __init__(self, used: int = 0) -> None:
self.used = used
self.audit: list[dict[str, object]] = []
def count_llm_usage_since(self, user_id: int, source: str, since: str) -> int:
return self.used
def record_llm_usage(
self,
user_id: int,
feature: str,
source: str,
model: str,
status: str,
latency_ms: int,
**metadata,
) -> None:
self.audit.append(
{
"user_id": user_id,
"feature": feature,
"source": source,
"model": model,
"status": status,
**metadata,
}
)
def profile() -> dict[str, object]:
return {
"source": "platform",
"primary": {"api_key": "p", "base_url": "https://p", "model": "primary"},
"fallback": {"api_key": "f", "base_url": "https://f", "model": "fallback"},
}
class LLMGatewayTests(unittest.TestCase):
def gateway(self, database: FakeDatabase | None = None) -> LLMGateway:
database = database or FakeDatabase()
return LLMGateway(
database=database,
user_id_supplier=lambda: 7,
membership_supplier=lambda: {"active": True},
settings_supplier=lambda: {"member_daily_limit": 50},
profile_supplier=profile,
)
def test_non_streaming_call_falls_back_and_audits_once(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
return {"answer": "ok"}
result = gateway.call("heaven_trend", "heaven-trend-v1", invoke, (ProviderFailure,))
self.assertEqual(result.role, "fallback")
self.assertEqual(result.value, {"answer": "ok"})
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "fallback")
self.assertEqual(database.audit[0]["prompt_version"], "heaven-trend-v1")
def test_stream_falls_back_before_first_delta(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
if model.role == "primary":
raise ProviderFailure("primary failed")
yield "a"
yield "b"
events = list(gateway.stream("mentor", "mentor-v1", invoke, (ProviderFailure,)))
self.assertEqual([event.value for event in events[:-1]], ["a", "b"])
self.assertEqual(events[-1].kind, "complete")
self.assertEqual(events[-1].role, "fallback")
self.assertEqual(database.audit[0]["status"], "success")
def test_stream_does_not_switch_model_after_output_started(self) -> None:
database = FakeDatabase()
gateway = self.gateway(database)
def invoke(model):
yield "first"
raise ProviderFailure(f"{model.role} interrupted")
iterator = gateway.stream("assistant", "assistant-v1", invoke, (ProviderFailure,))
self.assertEqual(next(iterator).value, "first")
with self.assertRaisesRegex(LLMGatewayError, "连接中断"):
list(iterator)
self.assertEqual(len(database.audit), 1)
self.assertEqual(database.audit[0]["model"], "primary")
self.assertEqual(database.audit[0]["status"], "failed")
def test_daily_quota_is_enforced_before_provider_call(self) -> None:
gateway = self.gateway(FakeDatabase(used=50))
with self.assertRaisesRegex(LLMGatewayError, "额度已用完"):
gateway.call("mentor", "mentor-v1", lambda model: "unused", (ProviderFailure,))
if __name__ == "__main__":
unittest.main()
+35
View File
@@ -0,0 +1,35 @@
from __future__ import annotations
import unittest
from llm_stream import OpenAIStreamAccumulator
class OpenAIStreamAccumulatorTests(unittest.TestCase):
def test_repeated_delta_chunks_are_preserved_as_model_output(self) -> None:
accumulator = OpenAIStreamAccumulator()
self.assertEqual(accumulator.feed({"delta": {"content": "yes"}}), "yes")
self.assertEqual(accumulator.feed({"delta": {"content": "yes"}}), "yes")
self.assertEqual(accumulator.text, "yesyes")
def test_final_snapshot_can_add_a_missing_suffix(self) -> None:
accumulator = OpenAIStreamAccumulator()
self.assertEqual(accumulator.feed({"delta": {"content": "first"}}), "first")
self.assertEqual(
accumulator.feed({"message": {"content": "first second"}}),
" second",
)
self.assertEqual(accumulator.text, "first second")
def test_incompatible_final_snapshot_is_not_appended_twice(self) -> None:
accumulator = OpenAIStreamAccumulator()
accumulator.feed({"delta": {"content": "streamed answer"}})
self.assertEqual(
accumulator.feed({"message": {"content": "rewritten final answer"}}),
"",
)
self.assertEqual(accumulator.text, "streamed answer")
if __name__ == "__main__":
unittest.main()
+294
View File
@@ -0,0 +1,294 @@
from __future__ import annotations
import tempfile
import unittest
from datetime import datetime, timedelta, timezone
from pathlib import Path
from database import ReviewDatabase
from market_insights import MarketInsightsService
from screener import FACTOR_FIELDS, ScreenerEngine
from tushare_client import TushareError
class FakeMarketClient:
def resolve_trade_context(self, requested: str):
value = str(requested).replace("-", "")
return value, "20260723"
def query(self, api_name, params=None, fields=""):
params = params or {}
date = params.get("trade_date", "")
if api_name == "trade_cal":
return [
{"cal_date": f"202607{day:02d}", "is_open": 1}
for day in range(14, 24)
]
if api_name == "stock_basic":
return [
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"},
{"ts_code": "000002.SZ", "name": "万科A", "industry": "房地产", "market": "主板", "list_date": "19910129"},
]
if api_name == "stk_auction":
if date == "20260724":
return []
return [
{"ts_code": "000001.SZ", "trade_date": date, "price": 10.5, "pre_close": 10, "vol": 20000, "amount": 5_000_000, "turnover_rate": 0.12, "volume_ratio": 1.8},
{"ts_code": "000002.SZ", "trade_date": date, "price": 9.8, "pre_close": 10, "vol": 10000, "amount": 2_000_000, "turnover_rate": 0.05, "volume_ratio": 0.8},
]
if api_name == "stk_limit":
return [
{"ts_code": "000001.SZ", "trade_date": date, "up_limit": 11, "down_limit": 9},
{"ts_code": "000002.SZ", "trade_date": date, "up_limit": 11, "down_limit": 9},
]
if api_name == "ths_index":
return [{"ts_code": "885001.TI", "name": "人工智能", "count": 2, "exchange": "A", "list_date": "20200101", "type": "N"}]
if api_name == "ths_daily":
if params.get("ts_code"):
return [
{"ts_code": "885001.TI", "trade_date": "20260722", "open": 99, "high": 102, "low": 98, "close": 101, "pct_change": 1, "vol": 100},
{"ts_code": "885001.TI", "trade_date": "20260723", "open": 101, "high": 104, "low": 100, "close": 103, "pct_change": 1.98, "vol": 120},
]
if date == "20260724":
return []
return [{"ts_code": "885001.TI", "trade_date": date, "close": 103, "pct_change": 1.98, "vol": 120, "turnover_rate": 2.3}]
if api_name == "ths_member":
return [
{"ts_code": "885001.TI", "con_code": "000001.SZ", "con_name": "平安银行"},
{"ts_code": "885001.TI", "con_code": "000002.SZ", "con_name": "万科A"},
]
if api_name == "daily":
return [
{"ts_code": "000001.SZ", "trade_date": date, "open": 10, "high": 11, "low": 9.8, "close": 10.5, "pct_chg": 5, "vol": 100, "amount": 200000},
{"ts_code": "000002.SZ", "trade_date": date, "open": 10, "high": 10, "low": 9.7, "close": 9.8, "pct_chg": -2, "vol": 100, "amount": 100000},
]
if api_name == "ths_hot":
if date == "20260724":
return []
return [
{"trade_date": date, "data_type": "热股", "ts_code": "000001.SZ", "ts_name": "平安银行", "rank": 1, "pct_change": 5, "current_price": 10.5, "hot": 1000, "concept": '["银行"]'},
{"trade_date": date, "data_type": "概念板块", "ts_code": "885001.TI", "ts_name": "人工智能", "rank": 1, "pct_change": 1.98, "hot": 800},
]
if api_name == "dc_hot":
if date == "20260724":
return []
return [{"trade_date": date, "data_type": "A股市场", "ts_code": "000001.SZ", "ts_name": "平安银行", "rank": 3, "pct_change": 5, "current_price": 10.5}]
return []
class ConfiguredIfind:
configured = True
class MarketInsightsTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.database = ReviewDatabase(Path(self.temp.name) / "review.db")
self.service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 20, tzinfo=timezone(timedelta(hours=8))
),
)
def tearDown(self):
self.temp.cleanup()
def test_auction_falls_back_and_normalizes_factors(self):
payload = self.service.auction_center("20260724")
self.assertEqual(payload["meta"]["trade_date"], "2026-07-23")
self.assertTrue(payload["meta"]["carried_forward"])
self.assertEqual(payload["meta"]["phase"], "observing")
self.assertEqual(payload["summary"]["stock_count"], 2)
self.assertEqual(payload["rows"][0]["amount_million"], 5)
self.assertEqual(payload["rows"][0]["change"], 5)
self.assertEqual(payload["rows"][0]["expectation"], "超预期")
self.assertEqual(set(payload["expectations"]), {"超预期", "符合预期", "低于预期"})
self.assertFalse(payload["news_feedback"]["available"])
self.assertEqual(len(payload["amount_history"]), 10)
self.assertEqual(payload["amount_history"][-1]["stock_count"], 2)
self.assertEqual(payload["focus_rows"][0]["code"], "000001")
def test_auction_amount_history_uses_the_same_a_share_universe_as_summary(self):
self.database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"},
{"ts_code": "688001.SH", "name": "首日上市", "industry": "半导体", "market": "科创板", "list_date": "20260723"},
])
self.database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": "20260723", "price": 10.5, "pre_close": 10, "amount": 5_000_000, "vol": 20_000},
{"ts_code": "688001.SH", "trade_date": "20260723", "price": 50, "pre_close": 10, "amount": 150_000_000, "vol": 3_000_000},
{"ts_code": "159001.SZ", "trade_date": "20260723", "price": 1.1, "pre_close": 1, "amount": 90_000_000, "vol": 90_000_000},
])
history = self.service._auction_amount_history("20260723")
self.assertEqual(history[-1]["stock_count"], 1)
self.assertEqual(history[-1]["amount_billion"], 0.05)
def test_real_limit_price_is_isolated_from_scored_candidates(self):
class OnePriceClient(FakeMarketClient):
def query(self, api_name, params=None, fields=""):
if api_name == "stk_limit":
return [
{"ts_code": "000001.SZ", "trade_date": "20260723", "up_limit": 10.5, "down_limit": 9},
{"ts_code": "000002.SZ", "trade_date": "20260723", "up_limit": 11, "down_limit": 9},
]
return super().query(api_name, params, fields)
service = MarketInsightsService(self.database, OnePriceClient(), self.service._now_provider)
payload = service.auction_center("20260724", force=True)
self.assertEqual([row["code"] for row in payload["one_price_rows"]], ["000001"])
self.assertNotIn("000001", {row["code"] for row in payload["rows"]})
def test_core_broken_pool_and_watchlist_are_kept_separate(self):
self.database.save_snapshot(
"20260723",
"test",
{
"limits": [{"code": "000001", "name": "平安银行", "sector": "银行", "streak": 3, "amount_billion": 8}],
"broken": [{"code": "000002", "name": "万科A", "sector": "房地产", "streak": 1}],
"sectors": [{"name": "银行", "count": 1, "leader": "平安银行"}],
},
)
user = self.database.create_user("auction-user", "salt", "hash")
self.database.save_watchlist(user["id"], "000002", "万科A", "房地产", "red")
payload = self.service.auction_center("20260724", force=True, user_id=user["id"])
rows = {row["code"]: row for row in payload["rows"]}
self.assertIn("昨日炸板", rows["000002"]["candidate_sources"])
self.assertIn("三板以上", rows["000001"]["core_tags"])
self.assertIn("000001", {row["code"] for row in payload["focus_rows"]})
self.assertEqual([row["code"] for row in payload["watchlist_rows"]], ["000002"])
anonymous = self.service.auction_center("20260724", user_id=0)
self.assertEqual(anonymous["watchlist_rows"], [])
def test_selection_window_does_not_disguise_previous_day_as_current(self):
service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 26, tzinfo=timezone(timedelta(hours=8))
),
)
payload = service.auction_center("20260724", force=True)
self.assertEqual(payload["meta"]["phase"], "selection")
self.assertFalse(payload["meta"]["available"])
self.assertFalse(payload["meta"]["carried_forward"])
self.assertEqual(payload["rows"], [])
def test_finalized_window_uses_and_persists_ifind_closing_snapshot(self):
service = MarketInsightsService(
self.database,
FakeMarketClient(),
now_provider=lambda: datetime(
2026, 7, 24, 9, 31, tzinfo=timezone(timedelta(hours=8))
),
ifind=ConfiguredIfind(),
)
calls = []
service._dynamic_auction_rows = lambda trade_date, baseline_date, user_id: (
calls.append((trade_date, baseline_date, user_id))
or [
{
"ts_code": "000001.SZ",
"trade_date": trade_date,
"price": 10.5,
"pre_close": 10,
"vol": 20_000,
"amount": 5_000_000,
"turnover_rate": 0.12,
"volume_ratio": 1.8,
"dynamic": True,
}
]
)
payload = service.auction_center("20260724")
cached = service.auction_center("20260724")
self.assertEqual(payload["meta"]["phase"], "finalized")
self.assertTrue(payload["meta"]["available"])
self.assertEqual(payload["summary"]["stock_count"], 1)
self.assertEqual(payload["rows"][0]["code"], "000001")
self.assertEqual(len(calls), 1)
self.assertTrue(cached["meta"]["cached"])
self.assertEqual(cached["summary"]["stock_count"], 1)
def test_theme_library_detail_and_popularity(self):
library = self.service.theme_library("20260724")
self.assertEqual(library["meta"]["trade_date"], "2026-07-23")
self.assertEqual(library["items"][0]["hot_rank"], 1)
detail = self.service.theme_detail("885001.TI", "20260724")
self.assertEqual(detail["summary"]["member_count"], 2)
self.assertEqual(detail["members"][0]["code"], "000001")
hot = self.service.popularity("20260724")
self.assertEqual(hot["summary"]["dual_count"], 1)
self.assertEqual(hot["combined"][0]["name"], "平安银行")
def test_feature_pages_use_local_data_when_trade_context_is_offline(self):
class OfflineClient:
def resolve_trade_context(self, requested: str):
raise TushareError("offline")
def query(self, api_name, params=None, fields=""):
raise TushareError("offline")
self.database.save_snapshot(
"20260723", "test", {"meta": {"trade_date": "2026-07-23"}}
)
self.database.save_snapshot(
"20260724", "test", {"meta": {"trade_date": "2026-07-24"}}
)
self.database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"}
])
self.database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": "20260724", "price": 10.5, "pre_close": 10, "amount": 5_000_000, "vol": 20_000, "turnover_rate": 0.12, "volume_ratio": 1.8}
])
self.database.save_data_snapshot(
"theme_library_v1", "20260724", "market",
{"meta": {"trade_date": "2026-07-24"}, "summary": {"theme_count": 1}, "items": [{"code": "885001.TI", "name": "人工智能", "member_count": 2}]},
)
self.database.save_data_snapshot(
"popularity_v1", "20260724", "market",
{"meta": {"trade_date": "2026-07-24"}, "summary": {"ths_count": 1, "dc_count": 0, "dual_count": 0}, "combined": [{"name": "平安银行"}], "ths": [], "dc": []},
)
service = MarketInsightsService(self.database, OfflineClient(), self.service._now_provider)
auction = service.auction_center("20260725")
self.assertEqual(auction["meta"]["trade_date"], "2026-07-24")
self.assertEqual(auction["summary"]["stock_count"], 1)
themes = service.theme_library("20260725")
self.assertTrue(themes["meta"]["cached"])
self.assertEqual(themes["items"][0]["name"], "人工智能")
popularity = service.popularity("20260725")
self.assertTrue(popularity["meta"]["cached"])
self.assertEqual(popularity["combined"][0]["name"], "平安银行")
class AuctionScreenerFactorTests(unittest.TestCase):
def test_auction_fields_are_available_to_formula_and_factor_rows(self):
with tempfile.TemporaryDirectory() as temporary:
database = ReviewDatabase(Path(temporary) / "review.db")
database.upsert_stock_master([
{"ts_code": "000001.SZ", "name": "平安银行", "industry": "银行", "market": "主板", "list_date": "19910403"}
])
dates = [f"202606{day:02d}" for day in range(1, 22)]
database.upsert_daily_bars([
{"ts_code": "000001.SZ", "trade_date": trade_date, "open": 10, "high": 11, "low": 9, "close": 10 + index * 0.1, "pct_chg": 1, "vol": 1000 + index, "amount": 200000}
for index, trade_date in enumerate(dates)
])
database.upsert_daily_indicators([
{"ts_code": "000001.SZ", "trade_date": dates[-1], "turnover_rate": 2, "volume_ratio": 1.2, "circ_mv": 100000, "total_mv": 120000}
])
database.upsert_auction_factors([
{"ts_code": "000001.SZ", "trade_date": dates[-1], "price": 12.6, "pre_close": 12, "amount": 8_000_000, "vol": 30000, "turnover_rate": 0.18, "volume_ratio": 2.1}
])
rows, actual_date = ScreenerEngine(database).build_factors(dates[-1])
self.assertEqual(actual_date, dates[-1])
self.assertEqual(rows[0]["auction_change"], 5)
self.assertEqual(rows[0]["auction_amount_million"], 8)
self.assertEqual(rows[0]["auction_volume_ratio"], 2.1)
self.assertIn("auction_change", FACTOR_FIELDS)
if __name__ == "__main__":
unittest.main()
+228
View File
@@ -0,0 +1,228 @@
from __future__ import annotations
import ast
import unittest
from datetime import datetime, timedelta, timezone
from pathlib import Path
from tushare_client import _sector_coverage_issue
def load_method(name: str):
source = Path("server.py").read_text(encoding="utf-8")
tree = ast.parse(source)
dashboard_service = next(
node for node in tree.body
if isinstance(node, ast.ClassDef) and node.name == "DashboardService"
)
method = next(
node for node in dashboard_service.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name == name
)
module = ast.Module(body=[method], type_ignores=[])
namespace = {
"datetime": datetime,
"Any": object,
"_sector_coverage_issue": _sector_coverage_issue,
}
exec(compile(ast.fix_missing_locations(module), "server.py", "exec"), namespace)
return namespace[name]
MARKET_MODE = load_method("_heaven_market_mode")
QUALITY_ISSUES = load_method("_heaven_trend_quality_issues")
TZ = timezone(timedelta(hours=8))
class MarketModeTests(unittest.TestCase):
def test_closed_rt_snapshot_is_closed_not_intraday(self):
dashboard = {"meta": {"realtime": True, "market_status": "closed"}}
now = datetime(2026, 7, 20, 16, 27, tzinfo=TZ)
self.assertEqual(MARKET_MODE("20260720", dashboard, now), "closed")
def test_trading_snapshot_is_intraday(self):
dashboard = {"meta": {"realtime": True, "market_status": "trading"}}
now = datetime(2026, 7, 20, 10, 30, tzinfo=TZ)
self.assertEqual(MARKET_MODE("20260720", dashboard, now), "intraday")
def test_historical_date_is_always_historical(self):
dashboard = {"meta": {"realtime": True, "market_status": "trading"}}
now = datetime(2026, 7, 20, 10, 30, tzinfo=TZ)
self.assertEqual(MARKET_MODE("20260717", dashboard, now), "historical")
@staticmethod
def intraday_layers(trade_date: str):
index_context = {
"trade_date": trade_date,
"precise": True,
"realtime": True,
"indices": [{"trade_date": trade_date}] * 3,
}
sector = {
"trade_date": trade_date,
"precise": True,
"realtime": True,
"taxonomy": "sw_l2",
"schema_version": 3,
"coverage": 100,
"relative_turnover": 1.2,
}
stock = {
"trade_date": trade_date,
"code": "002141",
"precise": True,
"realtime": True,
"turnover_source": "float_share",
"activity_source": "historical_progress",
}
return index_context, sector, stock
@staticmethod
def historical_layers(trade_date: str):
index_context = {
"trade_date": trade_date,
"precise": True,
"realtime": False,
"source": "tushare",
"indices": [{"trade_date": trade_date}] * 3,
}
sector = {
"trade_date": trade_date,
"precise": True,
"realtime": False,
"taxonomy": "sw_l2",
"schema_version": 3,
"source": "tushare_sw_daily+member_daily",
"coverage": 97,
}
stock = {
"trade_date": trade_date,
"code": "002141",
"precise": True,
"realtime": False,
"data_source": "tushare",
}
return index_context, sector, stock
def test_intraday_accepts_verified_realtime_layers(self):
trade_date = "20260720"
index_context, sector, stock = self.intraday_layers(trade_date)
dashboard = {
"meta": {
"realtime": True,
"market_status": "closed",
"updated_at": datetime.now(TZ).isoformat(),
}
}
issues = QUALITY_ISSUES(
trade_date, dashboard, index_context, sector, stock, "intraday"
)
self.assertEqual(issues, [])
def test_intraday_sector_coverage_90_passes_89_blocks(self):
trade_date = "20260720"
index_context, sector, stock = self.intraday_layers(trade_date)
dashboard = {"meta": {"realtime": True, "market_status": "closed"}}
sector["coverage"] = 90
self.assertEqual(
QUALITY_ISSUES(trade_date, dashboard, index_context, sector, stock, "intraday"),
[],
)
sector["coverage"] = 89
issues = QUALITY_ISSUES(
trade_date, dashboard, index_context, sector, stock, "intraday"
)
self.assertTrue(any("覆盖率" in issue for issue in issues))
def test_intraday_sector_requires_relative_turnover(self):
trade_date = "20260720"
index_context, sector, stock = self.intraday_layers(trade_date)
sector["relative_turnover"] = 0
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "intraday"
)
self.assertTrue(any("相对全市场换手" in issue for issue in issues))
def test_sector_must_use_shenwan_l2_taxonomy(self):
trade_date = "20260720"
index_context, sector, stock = self.intraday_layers(trade_date)
sector["taxonomy"] = "ths"
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "intraday"
)
self.assertTrue(any("申万二级" in issue for issue in issues))
def test_historical_accepts_official_daily_layers(self):
trade_date = "20260717"
index_context, sector, stock = self.historical_layers(trade_date)
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "historical"
)
self.assertEqual(issues, [])
def test_historical_rejects_realtime_index_layer(self):
trade_date = "20260717"
index_context, sector, stock = self.historical_layers(trade_date)
index_context["realtime"] = True
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "historical"
)
self.assertTrue(any("指数层" in issue for issue in issues))
def test_historical_rejects_realtime_sector_layer(self):
trade_date = "20260717"
index_context, sector, stock = self.historical_layers(trade_date)
sector["realtime"] = True
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "historical"
)
self.assertTrue(any("行业层" in issue for issue in issues))
def test_historical_rejects_non_tushare_stock(self):
trade_date = "20260717"
index_context, sector, stock = self.historical_layers(trade_date)
stock["data_source"] = "dashboard"
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "historical"
)
self.assertTrue(any("个股层" in issue for issue in issues))
def test_stock_must_be_precise(self):
trade_date = "20260717"
index_context, sector, stock = self.historical_layers(trade_date)
stock["precise"] = False
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "historical"
)
self.assertTrue(any("个股层" in issue for issue in issues))
def test_closed_mode_ignores_nonessential_dashboard_status(self):
trade_date = "20260720"
index_context, sector, stock = self.historical_layers(trade_date)
dashboard = {"meta": {"realtime": True, "market_status": "trading"}}
issues = QUALITY_ISSUES(
trade_date, dashboard, index_context, sector, stock, "closed"
)
self.assertEqual(issues, [])
def test_closed_mode_accepts_finalized_realtime_shenwan_snapshot(self):
trade_date = "20260720"
index_context, sector, stock = self.historical_layers(trade_date)
sector.update({
"realtime": True,
"finalized": True,
"inner_precise": True,
"outer_precise": True,
"relative_turnover": 1.2,
})
issues = QUALITY_ISSUES(
trade_date, {"meta": {}}, index_context, sector, stock, "closed"
)
self.assertEqual(issues, [])
if __name__ == "__main__":
unittest.main()
+93
View File
@@ -0,0 +1,93 @@
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from mentor_agent import MentorSkillRegistry
ROOT = Path(__file__).resolve().parents[1]
def write_skill(root: Path, directory: str, skill_id: str, name: str) -> None:
path = root / directory
path.mkdir(parents=True, exist_ok=True)
(path / "SKILL.md").write_text(
"\n".join(
(
"---",
f"name: {skill_id}",
"description: |",
f" {name}的思维框架。",
f" 用途:使用{name}的视角分析市场。",
"---",
"",
f"# {name} · 思维操作系统",
"",
'> "先看事实,再做判断。"',
"",
"### 模型1: 证据优先",
)
),
encoding="utf-8",
)
class MentorSkillRegistryTests(unittest.TestCase):
def test_private_skills_require_explicit_inclusion(self):
with tempfile.TemporaryDirectory() as temporary:
base = Path(temporary)
public_root = base / "public"
private_root = base / "private"
write_skill(public_root, "public-person", "public-person", "公开老师")
write_skill(private_root, "private-person", "private-person", "私有老师")
(public_root / "mentor_catalog.json").write_text(
json.dumps(
{
"mentors": {
"public-person": {
"evidence": {"grade": "A", "label": "原始语料", "note": "原帖"},
"quality": {"score": 6, "total": 6, "status": "pass"},
}
}
},
ensure_ascii=False,
),
encoding="utf-8",
)
registry = MentorSkillRegistry(public_root, private_root)
self.assertEqual([item.skill_id for item in registry.list_skills()], ["public-person"])
with self.assertRaises(ValueError):
registry.get_skill("private-person")
admin_skills = registry.list_skills(include_private=True)
self.assertEqual({item.skill_id for item in admin_skills}, {"public-person", "private-person"})
private_skill = registry.get_skill("private-person", include_private=True)
self.assertTrue(private_skill.is_private)
self.assertTrue(private_skill.public()["private"])
public_skill = registry.get_skill("public-person")
self.assertEqual(public_skill.evidence_grade, "A")
self.assertEqual(public_skill.quality_score, 6)
def test_all_public_project_skills_have_catalog_metadata(self):
registry = MentorSkillRegistry(ROOT / "游资skills")
skills = registry.list_skills()
self.assertGreaterEqual(len(skills), 21)
self.assertNotIn("xiaobai-perspective", {item.skill_id for item in skills})
self.assertTrue(all(item.evidence_grade in {"A", "B", "C"} for item in skills))
self.assertTrue(all(item.quality_total == 6 for item in skills))
def test_server_applies_private_guard_to_every_mentor_entry_point(self):
source = (ROOT / "server.py").read_text(encoding="utf-8")
mentor_section = source[source.index(" def mentor_setup"):source.index(" def _heaven_manual_schema")]
self.assertGreaterEqual(
mentor_section.count('include_private=self.membership()["is_admin"]'),
4,
)
if __name__ == "__main__":
unittest.main()
+108
View File
@@ -0,0 +1,108 @@
from __future__ import annotations
import json
import unittest
from pathlib import Path
from unittest.mock import patch
from mentor_agent import MentorSkill, chat_with_mentor, stream_with_mentor
class FakeStreamResponse:
def __init__(self, lines: list[bytes]) -> None:
self.lines = lines
def __enter__(self):
return iter(self.lines)
def __exit__(self, exc_type, exc_value, traceback):
return False
class MentorStreamTests(unittest.TestCase):
def setUp(self) -> None:
self.skill = MentorSkill(
skill_id="test-mentor",
name="测试老师",
description="测试",
tagline="先看事实",
focus=("纪律",),
content="只做条件化判断。",
path=Path("SKILL.md"),
)
self.lines = [
b'data: {"choices":[{"delta":{"content":"first"}}]}\n',
b'data: {"choices":[{"delta":{"content":" second"}}]}\n',
b"data: [DONE]\n",
]
def test_stream_requests_upstream_streaming_and_yields_deltas(self):
captured = {}
def open_request(request, timeout):
captured["payload"] = json.loads(request.data.decode("utf-8"))
captured["accept"] = request.headers.get("Accept")
return FakeStreamResponse(self.lines)
with patch("mentor_agent.urllib.request.urlopen", side_effect=open_request):
chunks = list(
stream_with_mentor(
self.skill, {"data_trade_date": "20260723"}, "怎么看?", [],
"key", "https://example.test/v1", "model",
)
)
self.assertEqual(chunks, ["first", " second"])
self.assertTrue(captured["payload"]["stream"])
self.assertEqual(captured["accept"], "text/event-stream")
def test_non_streaming_compatibility_wrapper_collects_chunks(self):
with patch(
"mentor_agent.urllib.request.urlopen",
return_value=FakeStreamResponse(self.lines),
):
result = chat_with_mentor(
self.skill, {}, "怎么看?", [], "key", "https://example.test/v1", "model"
)
self.assertEqual(result["answer"], "first second")
def test_final_full_message_does_not_duplicate_streamed_deltas(self):
lines = [
b'data: {"choices":[{"delta":{"content":"first"}}]}\n',
b'data: {"choices":[{"delta":{"content":" second"}}]}\n',
b'data: {"choices":[{"message":{"content":"first second"}}]}\n',
b"data: [DONE]\n",
]
with patch(
"mentor_agent.urllib.request.urlopen",
return_value=FakeStreamResponse(lines),
):
chunks = list(
stream_with_mentor(
self.skill, {}, "question", [], "key",
"https://example.test/v1", "model",
)
)
self.assertEqual(chunks, ["first", " second"])
def test_snapshot_only_stream_emits_only_new_suffix(self):
lines = [
b'data: {"choices":[{"message":{"content":"first"}}]}\n',
b'data: {"choices":[{"message":{"content":"first second"}}]}\n',
b"data: [DONE]\n",
]
with patch(
"mentor_agent.urllib.request.urlopen",
return_value=FakeStreamResponse(lines),
):
chunks = list(
stream_with_mentor(
self.skill, {}, "question", [], "key",
"https://example.test/v1", "model",
)
)
self.assertEqual(chunks, ["first", " second"])
if __name__ == "__main__":
unittest.main()
+116
View File
@@ -0,0 +1,116 @@
from __future__ import annotations
import unittest
from tushare_client import TushareClient
class FakeRealtimeClient(TushareClient):
def query(self, api_name, params=None, fields=""):
params = params or {}
if api_name == "trade_cal":
return [
{
"cal_date": "20260720",
"is_open": 1,
"pretrade_date": "20260717",
}
]
if api_name == "stock_basic":
return [
{"ts_code": "000001.SZ", "name": "", "industry": "银行"},
{"ts_code": "000002.SZ", "name": "", "industry": "地产"},
{"ts_code": "000003.SZ", "name": "", "industry": "元器件"},
]
if api_name == "stk_limit":
return [
{"ts_code": "000001.SZ", "up_limit": 11.0, "down_limit": 9.0},
{"ts_code": "000002.SZ", "up_limit": 22.0, "down_limit": 18.0},
{"ts_code": "000003.SZ", "up_limit": 33.0, "down_limit": 27.0},
]
if api_name == "daily_basic":
requested_code = str(params.get("ts_code") or "")
rows = [
{"ts_code": "000001.SZ", "trade_date": "20260717", "float_share": 1000},
{"ts_code": "000002.SZ", "trade_date": "20260717", "float_share": 2000},
{"ts_code": "000003.SZ", "trade_date": "20260717", "float_share": 3000},
]
return [row for row in rows if not requested_code or row["ts_code"] == requested_code]
if api_name == "daily":
return [
{"ts_code": params.get("ts_code"), "trade_date": "20260713", "vol": 1000, "amount": 1},
{"ts_code": params.get("ts_code"), "trade_date": "20260714", "vol": 1000, "amount": 1},
{"ts_code": params.get("ts_code"), "trade_date": "20260715", "vol": 1000, "amount": 1},
{"ts_code": params.get("ts_code"), "trade_date": "20260716", "vol": 1000, "amount": 1},
{"ts_code": params.get("ts_code"), "trade_date": "20260717", "vol": 1000, "amount": 1},
]
if api_name == "limit_list_d":
return [
{
"ts_code": "000001.SZ",
"name": "",
"industry": "银行",
"close": 10.0,
"pct_chg": 10.0,
"amount": 100000000,
"limit_times": 2,
}
]
if api_name == "rt_k":
rows = [
{
"ts_code": "000001.SZ", "name": "", "pre_close": 10.0,
"open": 10.1, "high": 11.0, "low": 10.0, "close": 11.0,
"vol": 1000, "amount": 100000000, "num": 10,
},
{
"ts_code": "000002.SZ", "name": "", "pre_close": 20.0,
"open": 19.5, "high": 20.0, "low": 18.0, "close": 18.0,
"vol": 2000, "amount": 200000000, "num": 20,
},
{
"ts_code": "000003.SZ", "name": "", "pre_close": 30.0,
"open": 31.0, "high": 33.0, "low": 30.0, "close": 32.0,
"vol": 3000, "amount": 300000000, "num": 30,
},
]
requested = {
code for code in str(params.get("ts_code") or "").split(",") if code
}
return [row for row in rows if row["ts_code"] in requested]
raise AssertionError(f"Unexpected API call: {api_name} {params}")
class RealtimeDashboardTests(unittest.TestCase):
def setUp(self):
TushareClient._realtime_reference_cache.clear()
TushareClient._capital_cache.clear()
TushareClient._latest_realtime_market.clear()
TushareClient._stock_activity_cache.clear()
self.client = FakeRealtimeClient("test-token")
def test_realtime_dashboard_classifies_pools_and_units(self):
dashboard = self.client._realtime_dashboard("20260720", "20260720", "20260717")
self.assertTrue(dashboard["meta"]["realtime"])
self.assertEqual(dashboard["meta"]["quote_count"], 3)
self.assertEqual(dashboard["overview"]["limit_up_count"], 1)
self.assertEqual(dashboard["overview"]["limit_down_count"], 1)
self.assertEqual(dashboard["overview"]["broken_count"], 1)
self.assertEqual(dashboard["overview"]["amount_billion"], 6.0)
self.assertEqual(dashboard["limits"][0]["streak"], 3)
self.assertEqual(dashboard["limits"][0]["amount_billion"], 1.0)
def test_realtime_stock_quote_uses_cached_industry(self):
self.client._load_realtime_reference("20260720", "20260717")
quote = self.client.realtime_stock_quote("000003.SZ")
self.assertEqual(quote["name"], "")
self.assertEqual(quote["sector"], "元器件")
self.assertAlmostEqual(quote["change"], 6.6667)
self.assertEqual(quote["amount_billion"], 3.0)
self.assertAlmostEqual(quote["turnover_rate"], 0.01)
if __name__ == "__main__":
unittest.main()
+54
View File
@@ -0,0 +1,54 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from alert_service import AlertService
from backend.bootstrap.container import build_application_container
from backend.database.repositories import (
SQLiteAlertRepository,
SQLiteStrategyTrackingRepository,
SQLiteTradeJournalRepository,
)
from database import ReviewDatabase
class RepositoryBoundaryTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.addCleanup(self.temporary.cleanup)
self.database = ReviewDatabase(Path(self.temporary.name) / "review.db")
def test_container_injects_narrow_user_repositories(self) -> None:
skills = Path(self.temporary.name) / "skills"
private = Path(self.temporary.name) / "private"
skills.mkdir()
private.mkdir()
container = build_application_container(self.database, {}, skills, private)
self.assertIsInstance(container.repositories.alerts, SQLiteAlertRepository)
self.assertIsInstance(container.repositories.trades, SQLiteTradeJournalRepository)
self.assertIsInstance(
container.repositories.strategy_tracking, SQLiteStrategyTrackingRepository
)
self.assertIs(container.alert_service.repository, container.repositories.alerts)
self.assertIs(container.trade_journal.repository, container.repositories.trades)
def test_private_repositories_reject_missing_account_owner(self) -> None:
alerts = SQLiteAlertRepository(self.database)
trades = SQLiteTradeJournalRepository(self.database)
tracking = SQLiteStrategyTrackingRepository(self.database)
with self.assertRaises(ValueError):
alerts.list_alerts(0, "20260729")
with self.assertRaises(ValueError):
trades.list_trade_entries(0)
with self.assertRaises(ValueError):
tracking.list_strategy_tracks(0)
def test_alert_service_keeps_database_compatible_structural_port(self) -> None:
service = AlertService(self.database)
self.assertIs(service.repository, self.database)
if __name__ == "__main__":
unittest.main()
+97
View File
@@ -0,0 +1,97 @@
from __future__ import annotations
import io
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from assistant_agent import ReviewAssistantError, stream_review_assistant
from database import ReviewDatabase
from server import RequestHandler
class StreamingResponse:
def __init__(self, lines: list[bytes]):
self.lines = lines
def __enter__(self):
return iter(self.lines)
def __exit__(self, exc_type, exc, traceback):
return False
class ReviewAssistantStreamingTests(unittest.TestCase):
def test_openai_compatible_sse_is_yielded_in_order(self):
response = StreamingResponse(
[
b'data: {"choices":[{"delta":{"content":"first"}}]}\n',
b'data: {"choices":[{"delta":{"content":" second"}}]}\n',
b'data: [DONE]\n',
]
)
with patch("assistant_agent.urllib.request.urlopen", return_value=response):
chunks = list(
stream_review_assistant({}, "question", [], "key", "https://example.test/v1", "model")
)
self.assertEqual(chunks, ["first", " second"])
def test_empty_stream_is_rejected(self):
response = StreamingResponse([b"data: [DONE]\n"])
with patch("assistant_agent.urllib.request.urlopen", return_value=response):
with self.assertRaises(ReviewAssistantError):
list(
stream_review_assistant({}, "question", [], "key", "https://example.test/v1", "model")
)
def test_final_full_message_does_not_duplicate_streamed_deltas(self):
response = StreamingResponse(
[
b'data: {"choices":[{"delta":{"content":"first"}}]}\n',
b'data: {"choices":[{"delta":{"content":" second"}}]}\n',
b'data: {"choices":[{"message":{"content":"first second"}}]}\n',
b"data: [DONE]\n",
]
)
with patch("assistant_agent.urllib.request.urlopen", return_value=response):
chunks = list(
stream_review_assistant(
{}, "question", [], "key", "https://example.test/v1", "model"
)
)
self.assertEqual(chunks, ["first", " second"])
def test_ndjson_event_writer_flushes_complete_line(self):
handler = RequestHandler.__new__(RequestHandler)
handler.wfile = io.BytesIO()
RequestHandler._write_stream_event(handler, {"type": "delta", "content": "片段"})
self.assertEqual(
handler.wfile.getvalue().decode("utf-8"),
'{"type":"delta","content":"片段"}\n',
)
class ReviewAssistantHistoryTests(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("assistant_owner", "salt", "hash")
self.other = self.database.create_user("assistant_other", "salt", "hash")
def tearDown(self) -> None:
self.temp.cleanup()
def test_history_is_private_and_clear_is_scoped(self):
self.database.save_assistant_exchange(
self.owner["id"], "今天怎么看?", "先看承接。", "20260722"
)
owner_messages = self.database.list_assistant_messages(self.owner["id"])
self.assertEqual([item["role"] for item in owner_messages], ["user", "assistant"])
self.assertEqual(self.database.list_assistant_messages(self.other["id"]), [])
self.assertEqual(self.database.delete_assistant_messages(self.other["id"]), 0)
self.assertEqual(self.database.delete_assistant_messages(self.owner["id"]), 2)
if __name__ == "__main__":
unittest.main()
+44
View File
@@ -0,0 +1,44 @@
from __future__ import annotations
import unittest
from sentiment_engine import _adaptive_score, _confirmed_phase
class SentimentEngineTests(unittest.TestCase):
def test_adaptive_score_uses_latest_250_observations(self):
history = [0.0] * 50 + [100.0] * 250
self.assertEqual(_adaptive_score(50.0, 50.0, history), 12.5)
def test_ice_must_repair_before_fermentation(self):
phase, reason = _confirmed_phase(
{"phase": "冰点"},
score=70,
day_change=45,
systemic_health=70,
profit_score=70,
ecology_score=75,
phase_signal="发酵",
extreme_ice=False,
fermentation_signal_count=2,
)
self.assertEqual(phase, "修复")
self.assertIn("冰点后", reason)
def test_repair_requires_continuous_fermentation_confirmation(self):
phase, _ = _confirmed_phase(
{"phase": "修复"},
score=58,
day_change=5,
systemic_health=55,
profit_score=60,
ecology_score=65,
phase_signal="发酵",
extreme_ice=False,
fermentation_signal_count=1,
)
self.assertEqual(phase, "修复")
if __name__ == "__main__":
unittest.main()
+167
View File
@@ -0,0 +1,167 @@
from __future__ import annotations
import threading
import unittest
from datetime import datetime, timedelta
from unittest.mock import patch
from server import DashboardService
class DetailDatabaseStub:
@staticmethod
def list_watchlist(user_id):
return []
@staticmethod
def list_notes(user_id, code=""):
return []
class RealtimeClientStub:
quote_calls = 0
def __init__(self, token):
self.token = token
@staticmethod
def resolve_trade_context(requested_date):
return requested_date, requested_date
@classmethod
def realtime_stock_quote(cls, ts_code, reference_date=""):
cls.quote_calls += 1
return {
"name": "测试股票",
"sector": "测试行业",
"price": 9.8,
"change": -2.0,
"open": 10.1,
"high": 10.2,
"low": 9.7,
"volume": 123400,
"amount_billion": 1.25,
"turnover_rate": 3.5,
}
class FixedMarketDatetime(datetime):
fixed_now = datetime.now().astimezone().replace(hour=10, minute=30, second=0, microsecond=0)
@classmethod
def now(cls, tz=None):
return cls.fixed_now
class FixedPreopenDatetime(datetime):
fixed_now = datetime.now().astimezone().replace(hour=8, minute=45, second=0, microsecond=0)
@classmethod
def now(cls, tz=None):
return cls.fixed_now
class StockDetailRealtimeTests(unittest.TestCase):
def setUp(self):
self.service = DashboardService.__new__(DashboardService)
self.service._system_credentials = {"tushare_token": "test-token"}
self.service.database = DetailDatabaseStub()
self.service._request_context = threading.local()
self.service._request_context.user_id = 1
RealtimeClientStub.quote_calls = 0
def test_today_detail_merges_rt_quote_without_mutating_daily_cache(self):
today = FixedMarketDatetime.fixed_now.strftime("%Y%m%d")
yesterday = (FixedMarketDatetime.fixed_now - timedelta(days=1)).strftime("%Y-%m-%d")
cached = {
"meta": {"trade_date": today, "source": "tushare"},
"stock": {"code": "002141", "name": "旧名称", "price": 10, "change": 7.1},
"prices": [
{
"trade_date": yesterday,
"open": 9.5,
"high": 10.1,
"low": 9.4,
"close": 10,
"change": 7.1,
"volume": 100,
}
],
"moneyflow": {},
}
with patch("server.datetime", FixedMarketDatetime), patch(
"server.TushareClient", RealtimeClientStub
):
result = self.service._prepare_stock_detail(cached, "002141", today)
self.assertEqual(result["meta"]["trade_date"], FixedMarketDatetime.fixed_now.strftime("%Y-%m-%d"))
self.assertTrue(result["meta"]["realtime"])
self.assertEqual(result["stock"]["price"], 9.8)
self.assertEqual(result["stock"]["change"], -2.0)
self.assertEqual(result["prices"][-1]["change"], -2.0)
self.assertEqual(result["prices"][-1]["trade_date"], FixedMarketDatetime.fixed_now.strftime("%Y-%m-%d"))
self.assertEqual(cached["stock"]["change"], 7.1)
self.assertEqual(len(cached["prices"]), 1)
self.assertEqual(RealtimeClientStub.quote_calls, 1)
def test_historical_detail_never_requests_realtime_quote(self):
historical = (FixedMarketDatetime.fixed_now - timedelta(days=5)).strftime("%Y%m%d")
payload = {
"meta": {"trade_date": historical, "source": "tushare"},
"stock": {"code": "002141", "price": 10, "change": 1.2},
"prices": [{"trade_date": historical, "close": 10, "change": 1.2}],
}
with patch("server.datetime", FixedMarketDatetime), patch(
"server.TushareClient", RealtimeClientStub
):
result = self.service._prepare_stock_detail(payload, "002141", historical)
self.assertEqual(result["stock"]["change"], 1.2)
self.assertFalse(result["meta"].get("realtime", False))
self.assertEqual(RealtimeClientStub.quote_calls, 0)
def test_today_detail_keeps_last_traded_bar_before_market_open(self):
today = FixedPreopenDatetime.fixed_now.strftime("%Y%m%d")
today_display = FixedPreopenDatetime.fixed_now.strftime("%Y-%m-%d")
yesterday = (FixedPreopenDatetime.fixed_now - timedelta(days=1)).strftime("%Y-%m-%d")
payload = {
"meta": {"trade_date": yesterday, "source": "tushare"},
"stock": {"code": "002141", "price": 10, "change": 0},
"prices": [
{
"trade_date": yesterday,
"open": 9.8,
"high": 10.1,
"low": 9.7,
"close": 10,
"change": 1.2,
"volume": 100,
},
{
"trade_date": today_display,
"open": 10,
"high": 10,
"low": 10,
"close": 10,
"change": 0,
"volume": 0,
"amount_billion": 0,
"realtime": True,
},
],
}
with patch("server.datetime", FixedPreopenDatetime), patch(
"server.TushareClient", RealtimeClientStub
):
result = self.service._prepare_stock_detail(payload, "002141", today)
self.assertEqual(result["meta"]["trade_date"], yesterday)
self.assertFalse(result["meta"].get("realtime", False))
self.assertEqual(result["prices"][-1]["trade_date"], yesterday)
self.assertEqual(result["stock"]["change"], 1.2)
self.assertEqual(RealtimeClientStub.quote_calls, 0)
if __name__ == "__main__":
unittest.main()
+158
View File
@@ -0,0 +1,158 @@
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()
+90
View File
@@ -0,0 +1,90 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from database import ReviewDatabase
from trade_journal import TradeJournalService
class TradeJournalTests(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("trade_owner", "salt", "hash")
self.other = self.database.create_user("trade_other", "salt", "hash")
self.service = TradeJournalService(self.database)
def tearDown(self) -> None:
self.temp.cleanup()
@staticmethod
def payload(**overrides):
data = {
"trade_date": "2026-07-22",
"code": "002141",
"name": "贤丰控股",
"action": "buy",
"price": 10.25,
"quantity": 1000,
"position_pct": 20,
"pnl_amount": "",
"pnl_pct": "",
"thesis": "修复期低吸",
"execution": "按计划成交",
"emotion": "calm",
"tags": "计划内,低吸",
}
data.update(overrides)
return data
def test_unrealized_entry_does_not_enter_win_rate(self):
self.service.save(self.owner["id"], self.payload())
result = self.service.list_entries(self.owner["id"])
self.assertEqual(result["summary"]["total"], 1)
self.assertEqual(result["summary"]["realized"], 0)
self.assertIsNone(result["summary"]["win_rate"])
self.assertEqual(result["items"][0]["tags"], ["计划内", "低吸"])
def test_realized_entries_build_summary(self):
self.service.save(
self.owner["id"], self.payload(action="sell", pnl_amount=500, pnl_pct=5)
)
self.service.save(
self.owner["id"],
self.payload(code="600000", name="浦发银行", action="sell", pnl_amount=-200, pnl_pct=-2, position_pct=40),
)
result = self.service.list_entries(self.owner["id"])
self.assertEqual(result["summary"]["realized"], 2)
self.assertEqual(result["summary"]["win_rate"], 50.0)
self.assertEqual(result["summary"]["pnl_amount"], 300.0)
self.assertEqual(result["summary"]["average_position"], 30.0)
def test_update_and_delete_require_ownership(self):
trade_id = self.service.save(self.owner["id"], self.payload())
with self.assertRaises(ValueError):
self.service.save(
self.other["id"], self.payload(id=trade_id, thesis="越权修改")
)
self.assertFalse(self.database.delete_trade_entry(self.other["id"], trade_id))
self.service.save(
self.owner["id"], self.payload(id=trade_id, thesis="复盘后修订")
)
self.assertEqual(
self.service.list_entries(self.owner["id"])["items"][0]["thesis"],
"复盘后修订",
)
self.assertTrue(self.database.delete_trade_entry(self.owner["id"], trade_id))
self.assertEqual(self.service.list_entries(self.owner["id"])["summary"]["total"], 0)
def test_entries_are_isolated_between_accounts(self):
self.service.save(self.owner["id"], self.payload())
self.assertEqual(self.service.list_entries(self.other["id"])["items"], [])
if __name__ == "__main__":
unittest.main()