diff --git a/database.py b/database.py index 6755f6d..73bd4cd 100644 --- a/database.py +++ b/database.py @@ -219,24 +219,43 @@ class ReviewDatabase: CREATE TABLE IF NOT EXISTS screener_strategies ( id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER, 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 + updated_at TEXT NOT NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE ); CREATE TABLE IF NOT EXISTS screener_runs ( id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER, 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 + created_at TEXT NOT NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE ); + + CREATE TABLE IF NOT EXISTS mentor_messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL, + mentor_id TEXT NOT NULL, + trade_date TEXT NOT NULL, + role TEXT NOT NULL, + content TEXT NOT NULL, + meta TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL, + FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE + ); + + CREATE INDEX IF NOT EXISTS idx_mentor_messages_conversation + ON mentor_messages(user_id, mentor_id, trade_date, id DESC); """ ) user_columns = { @@ -309,6 +328,38 @@ class ReviewDatabase: ON review_notes(user_id, trade_date DESC, id DESC) """ ) + strategy_columns = { + str(row["name"]) for row in connection.execute("PRAGMA table_info(screener_strategies)") + } + if "user_id" not in strategy_columns: + connection.execute("ALTER TABLE screener_strategies ADD COLUMN user_id INTEGER") + run_columns = { + str(row["name"]) for row in connection.execute("PRAGMA table_info(screener_runs)") + } + if "user_id" not in run_columns: + connection.execute("ALTER TABLE screener_runs ADD COLUMN user_id INTEGER") + if first_user and first_user["id"]: + first_user_id = int(first_user["id"]) + connection.execute( + "UPDATE screener_strategies SET user_id = ? WHERE builtin = 0 AND user_id IS NULL", + (first_user_id,), + ) + connection.execute( + "UPDATE screener_runs SET user_id = ? WHERE user_id IS NULL", + (first_user_id,), + ) + connection.execute( + """ + CREATE INDEX IF NOT EXISTS idx_screener_strategies_user + ON screener_strategies(user_id, builtin, updated_at DESC) + """ + ) + connection.execute( + """ + CREATE INDEX IF NOT EXISTS idx_screener_runs_user_date + ON screener_runs(user_id, trade_date DESC, id DESC) + """ + ) def count_users(self) -> int: with self.connect() as connection: @@ -658,6 +709,33 @@ class ReviewDatabase: except json.JSONDecodeError: return None + def get_latest_data_snapshot( + self, + kind: str, + cache_key_prefix: str, + maximum_cache_key: str, + exclude_source: str = "", + ) -> dict[str, Any] | None: + source_clause = " AND source != ?" if exclude_source else "" + parameters: list[Any] = [kind, f"{cache_key_prefix}%", maximum_cache_key] + if exclude_source: + parameters.append(exclude_source) + with self.connect() as connection: + row = connection.execute( + f""" + SELECT payload FROM data_snapshots + WHERE kind = ? AND cache_key LIKE ? AND cache_key <= ?{source_clause} + ORDER BY cache_key DESC LIMIT 1 + """, + parameters, + ).fetchone() + if not row: + return None + try: + return json.loads(row["payload"]) + except json.JSONDecodeError: + return None + def save_data_snapshot( self, kind: str, cache_key: str, source: str, payload: dict[str, Any] ) -> None: @@ -1073,7 +1151,7 @@ class ReviewDatabase: return result def save_screener_strategy( - self, name: str, description: str, regimes: list[str], formula: dict[str, Any], + self, user_id: int | None, name: str, description: str, regimes: list[str], formula: dict[str, Any], builtin: bool = False, strategy_id: int | None = None, ) -> int: now = datetime.now().astimezone().isoformat(timespec="seconds") @@ -1081,31 +1159,50 @@ class ReviewDatabase: formula_json = json.dumps(formula, ensure_ascii=False, separators=(",", ":")) with self.connect() as connection: if strategy_id: - cursor = connection.execute( - """ - UPDATE screener_strategies SET name=?, description=?, regimes=?, formula=?, - builtin=?, updated_at=? WHERE id=? - """, - (name, description, regimes_json, formula_json, int(builtin), now, strategy_id), - ) + if builtin: + cursor = connection.execute( + """ + UPDATE screener_strategies SET name=?, description=?, regimes=?, formula=?, + builtin=1, user_id=NULL, updated_at=? WHERE id=? AND builtin=1 + """, + (name, description, regimes_json, formula_json, now, strategy_id), + ) + else: + cursor = connection.execute( + """ + UPDATE screener_strategies SET name=?, description=?, regimes=?, formula=?, + updated_at=? WHERE id=? AND builtin=0 AND user_id=? + """, + (name, description, regimes_json, formula_json, now, strategy_id, int(user_id or 0)), + ) if cursor.rowcount == 0: raise ValueError("选股策略不存在。") return strategy_id cursor = connection.execute( """ INSERT INTO screener_strategies - (name, description, regimes, formula, builtin, created_at, updated_at) - VALUES (?, ?, ?, ?, ?, ?, ?) + (user_id, name, description, regimes, formula, builtin, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) """, - (name, description, regimes_json, formula_json, int(builtin), now, now), + (None if builtin else int(user_id or 0), name, description, regimes_json, formula_json, int(builtin), now, now), ) return int(cursor.lastrowid) - def list_screener_strategies(self) -> list[dict[str, Any]]: + def list_screener_strategies(self, user_id: int | None = None) -> list[dict[str, Any]]: with self.connect() as connection: - rows = connection.execute( - "SELECT * FROM screener_strategies ORDER BY builtin DESC, updated_at DESC, id" - ).fetchall() + if user_id is None: + rows = connection.execute( + "SELECT * FROM screener_strategies WHERE builtin = 1 ORDER BY updated_at DESC, id" + ).fetchall() + else: + rows = connection.execute( + """ + SELECT * FROM screener_strategies + WHERE builtin = 1 OR user_id = ? + ORDER BY builtin DESC, updated_at DESC, id + """, + (int(user_id),), + ).fetchall() result = [] for row in rows: item = dict(row) @@ -1115,24 +1212,26 @@ class ReviewDatabase: result.append(item) return result - def delete_screener_strategy(self, strategy_id: int) -> bool: + def delete_screener_strategy(self, user_id: int, strategy_id: int) -> bool: with self.connect() as connection: row = connection.execute( - "SELECT builtin FROM screener_strategies WHERE id = ?", + "SELECT builtin, user_id FROM screener_strategies WHERE id = ?", (strategy_id,), ).fetchone() if not row: raise ValueError("选股策略不存在。") if bool(row["builtin"]): raise ValueError("内置策略不能删除。") + if int(row["user_id"] or 0) != int(user_id): + raise ValueError("无权删除其他账号的策略。") cursor = connection.execute( - "DELETE FROM screener_strategies WHERE id = ? AND builtin = 0", - (strategy_id,), + "DELETE FROM screener_strategies WHERE id = ? AND builtin = 0 AND user_id = ?", + (strategy_id, int(user_id)), ) return cursor.rowcount > 0 def save_screener_run( - self, trade_date: str, regime: str, strategy_name: str, + self, user_id: int, trade_date: str, regime: str, strategy_name: str, formula: dict[str, Any], result: dict[str, Any], ) -> int: now = datetime.now().astimezone().isoformat(timespec="seconds") @@ -1140,23 +1239,23 @@ class ReviewDatabase: cursor = connection.execute( """ INSERT INTO screener_runs - (trade_date, regime, strategy_name, formula, result, created_at) - VALUES (?, ?, ?, ?, ?, ?) + (user_id, trade_date, regime, strategy_name, formula, result, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?) """, - (trade_date, regime, strategy_name, + (int(user_id), trade_date, regime, strategy_name, json.dumps(formula, ensure_ascii=False, separators=(",", ":")), json.dumps(result, ensure_ascii=False, separators=(",", ":")), now), ) return int(cursor.lastrowid) - def latest_screener_run(self, trade_date: str) -> dict[str, Any] | None: + def latest_screener_run(self, user_id: int, trade_date: str) -> dict[str, Any] | None: with self.connect() as connection: row = connection.execute( """ SELECT id, trade_date, regime, strategy_name, result, created_at - FROM screener_runs WHERE trade_date <= ? ORDER BY id DESC LIMIT 1 + FROM screener_runs WHERE user_id = ? AND trade_date <= ? ORDER BY id DESC LIMIT 1 """, - (trade_date,), + (int(user_id), trade_date), ).fetchone() if not row: return None @@ -1168,6 +1267,60 @@ class ReviewDatabase: result["meta"]["created_at"] = row["created_at"] return result + def save_mentor_exchange( + self, + user_id: int, + mentor_id: str, + trade_date: str, + question: str, + answer: str, + meta: str = "", + ) -> None: + now = datetime.now().astimezone().isoformat(timespec="seconds") + with self.connect() as connection: + connection.executemany( + """ + INSERT INTO mentor_messages + (user_id, mentor_id, trade_date, role, content, meta, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?) + """, + [ + (int(user_id), mentor_id, trade_date, "user", question, "", now), + (int(user_id), mentor_id, trade_date, "assistant", answer, meta, now), + ], + ) + connection.execute( + """ + DELETE FROM mentor_messages + WHERE user_id = ? AND id NOT IN ( + SELECT id FROM mentor_messages WHERE user_id = ? ORDER BY id DESC LIMIT 500 + ) + """, + (int(user_id), int(user_id)), + ) + + def list_mentor_messages( + self, user_id: int, mentor_id: str, trade_date: str, limit: int = 100 + ) -> list[dict[str, Any]]: + with self.connect() as connection: + rows = connection.execute( + """ + SELECT role, content, meta, created_at FROM mentor_messages + WHERE user_id = ? AND mentor_id = ? AND trade_date = ? + ORDER BY id DESC LIMIT ? + """, + (int(user_id), mentor_id, trade_date, max(1, min(500, int(limit)))), + ).fetchall() + return [dict(row) for row in reversed(rows)] + + def delete_mentor_messages(self, user_id: int, mentor_id: str, trade_date: str) -> int: + with self.connect() as connection: + cursor = connection.execute( + "DELETE FROM mentor_messages WHERE user_id = ? AND mentor_id = ? AND trade_date = ?", + (int(user_id), mentor_id, trade_date), + ) + return int(cursor.rowcount) + def start_sync(self, trade_date: str, source: str) -> int: started_at = datetime.now().astimezone().isoformat(timespec="seconds") with self.connect() as connection: diff --git a/screener.py b/screener.py index e730f56..6d58b31 100644 --- a/screener.py +++ b/screener.py @@ -259,7 +259,7 @@ class ScreenerEngine: existing = {item["name"] for item in self.database.list_screener_strategies() if item["builtin"]} for strategy in BUILTIN_STRATEGIES: if strategy["name"] not in existing: - self.database.save_screener_strategy(**strategy, builtin=True) + self.database.save_screener_strategy(None, **strategy, builtin=True) def detect_regime(self, trade_date: str) -> dict[str, Any]: series = latest_contiguous_history( @@ -336,7 +336,7 @@ class ScreenerEngine: return result def screen( - self, trade_date: str, formula: dict[str, Any], regime: str, + self, user_id: int, trade_date: str, formula: dict[str, Any], regime: str, strategy_name: str, run_backtest: bool = True, realtime_snapshot: dict[str, Any] | None = None, ) -> dict[str, Any]: @@ -384,7 +384,7 @@ class ScreenerEngine: "disclaimer": "概率为历史条件估计,不代表未来收益;退潮或样本不足时允许无候选。", } run_id = self.database.save_screener_run( - actual_date, regime, strategy_name, formula, result + user_id, actual_date, regime, strategy_name, formula, result ) result["meta"]["run_id"] = run_id return result diff --git a/server.py b/server.py index 593f789..a530b0e 100644 --- a/server.py +++ b/server.py @@ -18,7 +18,6 @@ from typing import Any from urllib.parse import parse_qs, unquote, urlparse from database import ReviewDatabase -from demo_data import build_demo_dragon_tiger, build_demo_stock_detail from heaven_agent import HeavenAgentError, interpret_heaven from heaven_engine import ( _market_line_scores, @@ -1074,7 +1073,7 @@ class DashboardService: llm_access = self.llm_access_status() return { "configured": self.configured, - "mode": "tushare" if self.configured else "demo", + "mode": "tushare" if self.configured else "unavailable", "llm_configured": self.llm_configured, "llm_model": self.llm_primary_model if self.llm_configured else "", "llm_fallback_configured": self.llm_fallback_configured, @@ -1097,7 +1096,7 @@ class DashboardService: "trade_date": normalized_date, "regime": regime, "regimes": [{"id": key, "label": value} for key, value in REGIMES.items()], - "strategies": self.database.list_screener_strategies(), + "strategies": self.database.list_screener_strategies(self.current_user_id), "factor_fields": [{"id": key, "label": value} for key, value in FACTOR_FIELDS.items()], "factor_data": { "date_count": len(factor_dates), @@ -1111,7 +1110,7 @@ class DashboardService: "fallback_configured": self.llm_fallback_configured, "fallback_model": self.llm_fallback_model if self.llm_fallback_configured else "", }, - "latest_result": self.database.latest_screener_run(normalized_date), + "latest_result": self.database.latest_screener_run(self.current_user_id, normalized_date), } def sync_screener_data(self, trade_date: str, lookback: int = 45) -> dict[str, Any]: @@ -1191,14 +1190,19 @@ class DashboardService: if not isinstance(regimes, list) or not regimes or any(item not in REGIMES for item in regimes): raise ValueError("策略适用阶段不正确。") formula = self.screener.validate_formula(payload.get("formula") or {}) - strategy_id = self.database.save_screener_strategy(name, description, regimes, formula) - return {"id": strategy_id, "strategies": self.database.list_screener_strategies()} + strategy_id = self.database.save_screener_strategy( + self.current_user_id, name, description, regimes, formula + ) + return { + "id": strategy_id, + "strategies": self.database.list_screener_strategies(self.current_user_id), + } def delete_screener_strategy(self, strategy_id: int) -> dict[str, Any]: - deleted = self.database.delete_screener_strategy(strategy_id) + deleted = self.database.delete_screener_strategy(self.current_user_id, strategy_id) return { "deleted": deleted, - "strategies": self.database.list_screener_strategies(), + "strategies": self.database.list_screener_strategies(self.current_user_id), } def mentor_setup(self, trade_date: str) -> dict[str, Any]: @@ -1275,6 +1279,14 @@ class DashboardService: "success", int(result.get("latency_ms") or 0), ) + self.database.save_mentor_exchange( + self.current_user_id, + mentor_id, + trade_date, + question, + str(result.get("answer") or ""), + context["data_trade_date"], + ) return { **result, "mentor": skill.public(), @@ -1284,6 +1296,21 @@ class DashboardService: "notice": "主模型调用失败,已自动切换辅助模型。" if compiler == "fallback" else "", } + def mentor_messages(self, mentor_id: str, trade_date: str) -> list[dict[str, Any]]: + mentor_id = validate_text(mentor_id, "问师角色", 100, required=True) + trade_date = normalize_date(trade_date) + self.mentor_skills.get_skill(mentor_id) + return self.database.list_mentor_messages( + self.current_user_id, mentor_id, trade_date + ) + + def clear_mentor_messages(self, mentor_id: str, trade_date: str) -> int: + mentor_id = validate_text(mentor_id, "问师角色", 100, required=True) + trade_date = normalize_date(trade_date) + return self.database.delete_mentor_messages( + self.current_user_id, mentor_id, trade_date + ) + @staticmethod def _heaven_manual_schema(market_mode: str) -> dict[str, dict[str, Any]]: intraday = market_mode == "intraday" @@ -2585,7 +2612,7 @@ class DashboardService: except TushareError as exc: raise ValueError(f"实时选股行情不可用,已停止筛选:{exc}") from exc return self.screener.screen( - trade_date, formula, regime, strategy_name, + self.current_user_id, trade_date, formula, regime, strategy_name, bool(payload.get("run_backtest", True)), realtime_snapshot, ) @@ -2636,17 +2663,29 @@ class DashboardService: self.database.save_data_snapshot(cache_kind, normalized_date, "tushare", payload) return payload - demo = self._apply_seat_aliases(build_demo_dragon_tiger(normalized_date)) - demo["meta"] = { - **demo.get("meta", {}), + return { + "meta": { "requested_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", - "source": "demo", - "status": "demo", + "trade_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", + "source": "unavailable", + "status": "unavailable", "schema_version": 3, "cached": False, - "notice": "尚未配置 Tushare Token,当前展示演示数据。", + "notice": "公共行情尚未配置,暂无龙虎榜数据。", + }, + "summary": { + "trader_count": 0, + "identity_count": 0, + "operation_count": 0, + "active_stock_count": 0, + "seat_net_buy_million": 0, + "unclassified_count": 0, + "directory_count": 0, + }, + "traders": [], + "unclassified_seats": [], + "rows": [], } - return demo def _search_market_directory(self) -> list[dict[str, Any]]: cached = self.database.get_data_snapshot("search_directory", "ths") or {} @@ -2935,12 +2974,12 @@ class DashboardService: cache_key = f"{code}:{normalized_date}" if not force: cached = self.database.get_data_snapshot("stock_detail", cache_key) - if cached: + if cached and str((cached.get("meta") or {}).get("source") or "") != "demo": cached["meta"] = {**cached.get("meta", {}), "cached": True} return self._enrich_stock_detail(cached) name, sector = self._stock_identity(code, normalized_date) - source = "tushare" if self.configured else "demo" + source = "tushare" if self.configured: try: payload = TushareClient(self.token).stock_detail( @@ -2949,16 +2988,31 @@ class DashboardService: if not payload.get("prices"): raise TushareError("No price history returned") except TushareError as exc: - source = "demo" - payload = build_demo_stock_detail( - code, - normalized_date, - name, - sector, - f"个股行情接口暂不可用,已回退演示数据。原因:{exc}", + payload = self.database.get_latest_data_snapshot( + "stock_detail", f"{code}:", cache_key, exclude_source="demo" ) + if not payload: + raise ValueError(f"暂无 {code} 的真实行情数据:{exc}") from exc + payload = copy.deepcopy(payload) + payload["meta"] = { + **payload.get("meta", {}), + "cached": True, + "notice": "最新行情暂不可用,已沿用最近真实收盘数据。", + } + return self._enrich_stock_detail(payload) else: - payload = build_demo_stock_detail(code, normalized_date, name, sector) + payload = self.database.get_latest_data_snapshot( + "stock_detail", f"{code}:", cache_key, exclude_source="demo" + ) + if not payload: + raise ValueError(f"暂无 {code} 的真实行情数据,请等待后台完成首次同步。") + payload = copy.deepcopy(payload) + payload["meta"] = { + **payload.get("meta", {}), + "cached": True, + "notice": "公共行情尚未配置,已沿用最近真实收盘数据。", + } + return self._enrich_stock_detail(payload) payload["meta"]["source"] = source payload["meta"]["cached"] = False self.database.save_data_snapshot("stock_detail", cache_key, source, payload) @@ -3043,7 +3097,7 @@ class DashboardService: return { "meta": { "trade_date": resolved_date, - "source": detail_meta.get("source") or "demo", + "source": detail_meta.get("source") or "unavailable", "notice": detail_meta.get("notice") or "", "intraday_status": intraday_status, "intraday_notice": intraday_notice, @@ -3317,7 +3371,10 @@ class RequestHandler(BaseHTTPRequestHandler): return if parsed.path.startswith("/api/admin/") and not self.require_admin(): return - if parsed.path in {"/api/screener/setup", "/api/mentors/setup", "/api/heaven/setup"} and not self.require_member(): + if parsed.path in { + "/api/screener/setup", "/api/mentors/setup", "/api/mentors/messages", + "/api/heaven/setup", + } and not self.require_member(): return if parsed.path == "/api/admin/settings": self.send_json( @@ -3462,6 +3519,20 @@ class RequestHandler(BaseHTTPRequestHandler): except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return + if parsed.path == "/api/mentors/messages": + query = parse_qs(parsed.query) + try: + self.send_json( + { + "items": SERVICE.mentor_messages( + query.get("mentor_id", [""])[0], + query.get("trade_date", [date.today().isoformat()])[0], + ) + } + ) + except ValueError as exc: + self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) + return if parsed.path == "/api/heaven/setup": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] @@ -3537,12 +3608,18 @@ class RequestHandler(BaseHTTPRequestHandler): self.save_note() return if parsed.path == "/api/reasons": + if not self.require_admin(): + return self.save_reason() return if parsed.path == "/api/seat-aliases": + if not self.require_admin(): + return self.save_seat_alias() return if parsed.path == "/api/heaven/sector-phases": + if not self.require_admin(): + return self.save_sector_phase_override() return if parsed.path == "/api/backfill": @@ -3582,6 +3659,19 @@ class RequestHandler(BaseHTTPRequestHandler): deleted = SERVICE.database.delete_user_birth_profile(SERVICE.current_user_id) self.send_json({"ok": True, "deleted": deleted}) return + if parsed.path == "/api/mentors/messages": + if not self.require_member(): + return + query = parse_qs(parsed.query) + try: + deleted = SERVICE.clear_mentor_messages( + query.get("mentor_id", [""])[0], + query.get("trade_date", [date.today().isoformat()])[0], + ) + self.send_json({"ok": True, "deleted": deleted}) + except ValueError as exc: + self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) + return strategy_match = re.fullmatch(r"/api/screener/strategies/(\d+)", parsed.path) if strategy_match: if not self.require_member(): @@ -3608,7 +3698,7 @@ class RequestHandler(BaseHTTPRequestHandler): return sector_phase_match = re.fullmatch(r"/api/heaven/sector-phases/(.+)", parsed.path) if sector_phase_match: - if not self.require_member(): + if not self.require_admin(): return name = unquote(sector_phase_match.group(1)).strip() deleted = SERVICE.database.delete_sector_phase_override(name) diff --git a/static/app.js b/static/app.js index 5b92082..5d358b1 100644 --- a/static/app.js +++ b/static/app.js @@ -286,6 +286,8 @@ async function applyAuthenticatedSession(session) { updateAccountIdentityBadges(session.user?.membership || {}); document.querySelector("#settingsButton").hidden = !isAdmin; document.querySelector("#syncButton").hidden = !isAdmin; + document.querySelector("#reasonForm").hidden = !isAdmin; + document.querySelector("#sectorPhaseManager").hidden = !isAdmin; const factorSyncButton = document.querySelector("#factorSyncButton"); if (factorSyncButton) factorSyncButton.hidden = false; document.querySelector("#authGate").hidden = true; @@ -1309,8 +1311,9 @@ function renderDragonTraderDetail(trader) { function renderUnclassifiedSeats() { const seats = state.dragonTiger?.unclassified_seats || []; - document.querySelector("#dragonUnclassifiedSection").hidden = seats.length === 0; - document.querySelector("#dragonUnclassifiedFilter").hidden = seats.length === 0; + const canManage = state.user?.role === "admin"; + document.querySelector("#dragonUnclassifiedSection").hidden = !canManage || seats.length === 0; + document.querySelector("#dragonUnclassifiedFilter").hidden = !canManage || seats.length === 0; if (!seats.length && state.dragonFilter === "unclassified") { state.dragonFilter = "all"; document.querySelectorAll("[data-dragon-filter]").forEach((button) => { @@ -1826,7 +1829,7 @@ async function loadMentorSetup(force = false) { state.mentorSetup = payload; const selectedExists = payload.mentors.some((item) => item.id === state.selectedMentorId); state.selectedMentorId = selectedExists ? state.selectedMentorId : payload.mentors[0]?.id || ""; - state.mentorMessages = loadStoredMentorMessages(); + state.mentorMessages = await loadMentorMessages(); renderMentorWorkspace(); } catch (error) { showMentorNotice(error.message || "问师模块加载失败"); @@ -1854,13 +1857,14 @@ function renderMentorWorkspace() { renderMentorMessages(); } -function selectMentor(mentorId) { +async function selectMentor(mentorId) { if (mentorId === state.selectedMentorId) return; - saveStoredMentorMessages(); state.selectedMentorId = mentorId; - state.mentorMessages = loadStoredMentorMessages(); + state.mentorMessages = []; hideMentorNotice(); renderMentorWorkspace(); + state.mentorMessages = await loadMentorMessages(); + renderMentorMessages(); } function renderMentorMessages() { @@ -1925,7 +1929,6 @@ async function sendMentorQuestion(event) { content: payload.answer, meta: `${displayCompactDate(payload.data_trade_date)} · ${modelRole} ${payload.model} · ${number(payload.latency_ms)}ms`, }); - saveStoredMentorMessages(); if (payload.notice) showMentorNotice(payload.notice); setStatus("问师回答完成"); } catch (error) { @@ -1945,33 +1948,37 @@ function useMentorQuickPrompt(prompt) { input.focus(); } -function clearMentorConversation() { +async function clearMentorConversation() { if (!state.mentorMessages.length || !window.confirm("确定清空当前老师的对话记录吗?")) return; - state.mentorMessages = []; - try { localStorage.removeItem(mentorStorageKey()); } catch {} - hideMentorNotice(); - renderMentorMessages(); -} - -function mentorStorageKey() { - const dateKey = (state.mentorSetup?.trade_date || elements.tradeDate.value).replaceAll("-", ""); - return `xiaobai-mentor-chat:${state.selectedMentorId}:${dateKey}`; -} - -function loadStoredMentorMessages() { - if (!state.selectedMentorId) return []; try { - const messages = JSON.parse(localStorage.getItem(mentorStorageKey()) || "[]"); - if (!Array.isArray(messages)) return []; - return messages.filter((item) => ["user", "assistant"].includes(item?.role) && typeof item.content === "string").slice(-20); - } catch { - return []; + const query = new URLSearchParams({ + mentor_id: state.selectedMentorId, + trade_date: state.mentorSetup?.trade_date || elements.tradeDate.value, + }); + await apiRequest(`/api/mentors/messages?${query}`, "DELETE"); + state.mentorMessages = []; + hideMentorNotice(); + renderMentorMessages(); + } catch (error) { + showToast(error.message || "对话记录清空失败"); } } -function saveStoredMentorMessages() { - if (!state.selectedMentorId) return; - try { localStorage.setItem(mentorStorageKey(), JSON.stringify(state.mentorMessages.slice(-20))); } catch {} +async function loadMentorMessages() { + if (!state.selectedMentorId) return []; + try { + const query = new URLSearchParams({ + mentor_id: state.selectedMentorId, + trade_date: state.mentorSetup?.trade_date || elements.tradeDate.value, + }); + const payload = await apiRequest(`/api/mentors/messages?${query}`); + return (payload.items || []).filter( + (item) => ["user", "assistant"].includes(item?.role) && typeof item.content === "string", + ).slice(-100); + } catch (error) { + showMentorNotice(error.message || "对话记录加载失败"); + return []; + } } function showMentorNotice(message) { @@ -2744,11 +2751,12 @@ function drawQiUseConnections(animate = false) { function renderSectorPhaseOverrides(items) { const container = document.querySelector("#sectorPhaseOverrides"); + const canManage = state.user?.role === "admin"; container.innerHTML = items.length ? items.map((item) => `
${escapeHtml(item.element)} ${escapeHtml(item.name)} - + ${canManage ? `` : ""}
`).join("") : '

暂无手动归类

'; container.querySelectorAll("[data-sector-phase-delete]").forEach((button) => { diff --git a/static/index.html b/static/index.html index 0e4012e..1b1dd29 100644 --- a/static/index.html +++ b/static/index.html @@ -544,7 +544,7 @@

历法细目

中运、司天、在泉与节气定位
-
+
管理手动归类精确名称优先
@@ -1036,7 +1036,7 @@

事件逻辑

--

-- - +
diff --git a/tests/test_account_data_boundaries.py b/tests/test_account_data_boundaries.py new file mode 100644 index 0000000..85a08d5 --- /dev/null +++ b/tests/test_account_data_boundaries.py @@ -0,0 +1,193 @@ +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_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_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'); + """ + ) + 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")) + + +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") + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_frontend_contract.py b/tests/test_frontend_contract.py index 9925340..36612f1 100644 --- a/tests/test_frontend_contract.py +++ b/tests/test_frontend_contract.py @@ -44,6 +44,11 @@ class FrontendContractTests(unittest.TestCase): self.assertEqual(views, navigation) self.assertEqual(len(views), 13) + 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) + if __name__ == "__main__": unittest.main()