from __future__ import annotations import argparse import calendar import copy import json import mimetypes import os import re import secrets import threading from datetime import date, datetime, timedelta, timezone from http import HTTPStatus from http.cookies import SimpleCookie from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from typing import Any from urllib.parse import parse_qs, unquote, urlparse from database import ReviewDatabase from heaven_agent import HeavenAgentError, interpret_heaven from heaven_engine import ( _market_line_scores, _score_to_line, build_five_phase_field, build_market_hexagram, build_personal_field, hexagram_from_lines, ) from llm_strategy import LLMCompilerError, compile_strategy_with_llm, test_llm_connection from mentor_agent import MentorAgentError, MentorSkillRegistry, chat_with_mentor from realtime_aggregator import WebRealtimeAggregator from screener import ( FACTOR_FIELDS, REGIMES, FactorDataService, ScreenerEngine, compile_local_strategy, ) from security import SecretVault, hash_password, token_hash, verify_password from sentiment_engine import ( COMPONENT_WEIGHTS, apply_sentiment_to_dashboard, build_sentiment_history, latest_contiguous_history, ) from tushare_client import TushareClient, TushareError APP_DIR = Path(__file__).resolve().parent STATIC_DIR = APP_DIR / "static" DATA_DIR = APP_DIR / "data" ENV_FILE = APP_DIR / ".env" MENTOR_SKILLS_DIR = APP_DIR / "游资skills" TOKEN_PATTERN = re.compile(r"^[A-Za-z0-9_-]{20,128}$") USERNAME_PATTERN = re.compile(r"^[A-Za-z0-9_\-\u4e00-\u9fff]{3,30}$") SESSION_COOKIE = "xiaobai_session" SESSION_MAX_AGE = 30 * 24 * 60 * 60 LEGACY_SECRET_KEYS = { "TUSHARE_TOKEN", "LLM_API_KEY", "LLM_BASE_URL", "LLM_MODEL", "LLM_PRIMARY_API_KEY", "LLM_PRIMARY_BASE_URL", "LLM_PRIMARY_MODEL", "LLM_FALLBACK_API_KEY", "LLM_FALLBACK_BASE_URL", "LLM_FALLBACK_MODEL", } SEARCH_INDEXES = ( {"id": "000001.SH", "code": "000001.SH", "name": "上证指数", "type": "index", "subtitle": "沪市综合指数"}, {"id": "399001.SZ", "code": "399001.SZ", "name": "深证成指", "type": "index", "subtitle": "深市成份指数"}, {"id": "399006.SZ", "code": "399006.SZ", "name": "创业板指", "type": "index", "subtitle": "创业板核心指数"}, ) SEARCH_TYPE_LABELS = { "stock": "股票", "sector": "板块", "theme": "题材", "index": "指数", } THS_SEARCH_TYPES = { "I": ("sector", "行业板块"), "R": ("sector", "地域板块"), "N": ("theme", "概念题材"), } def load_local_env() -> None: if not ENV_FILE.exists(): return for raw_line in ENV_FILE.read_text(encoding="utf-8").splitlines(): line = raw_line.strip() if not line or line.startswith("#") or "=" not in line: continue key, value = line.split("=", 1) os.environ.setdefault(key.strip(), value.strip().strip('"').strip("'")) def save_local_env(updates: dict[str, str]) -> None: values: dict[str, str] = {} if ENV_FILE.exists(): for raw_line in ENV_FILE.read_text(encoding="utf-8").splitlines(): if "=" in raw_line and not raw_line.lstrip().startswith("#"): key, value = raw_line.split("=", 1) values[key.strip()] = value.strip().strip('"').strip("'") values.update(updates) ENV_FILE.write_text( "".join(f"{key}={value}\n" for key, value in values.items()), encoding="utf-8", ) def remove_local_env(keys: set[str]) -> None: if not ENV_FILE.exists(): return kept = [] for raw_line in ENV_FILE.read_text(encoding="utf-8").splitlines(): if "=" in raw_line and not raw_line.lstrip().startswith("#"): key = raw_line.split("=", 1)[0].strip() if key in keys: continue kept.append(raw_line) ENV_FILE.write_text("".join(f"{line}\n" for line in kept), encoding="utf-8") for key in keys: os.environ.pop(key, None) class DashboardService: def __init__(self) -> None: load_local_env() environment_credentials = { "tushare_token": os.environ.get("TUSHARE_TOKEN", "").strip(), "platform_llm_primary_api_key": os.environ.get( "LLM_PRIMARY_API_KEY", os.environ.get("LLM_API_KEY", "") ).strip(), "platform_llm_primary_base_url": os.environ.get( "LLM_PRIMARY_BASE_URL", os.environ.get("LLM_BASE_URL", "https://api.openai.com/v1") ).strip(), "platform_llm_primary_model": os.environ.get( "LLM_PRIMARY_MODEL", os.environ.get("LLM_MODEL", "") ).strip(), "platform_llm_fallback_api_key": os.environ.get("LLM_FALLBACK_API_KEY", "").strip(), "platform_llm_fallback_base_url": os.environ.get("LLM_FALLBACK_BASE_URL", "").strip(), "platform_llm_fallback_model": os.environ.get("LLM_FALLBACK_MODEL", "").strip(), } encryption_key = os.environ.get("APP_ENCRYPTION_KEY", "").strip() if not encryption_key: encryption_key = SecretVault.generate_key() save_local_env({"APP_ENCRYPTION_KEY": encryption_key}) os.environ["APP_ENCRYPTION_KEY"] = encryption_key self.vault = SecretVault(encryption_key) self.database = ReviewDatabase(DATA_DIR / "review.db") self.sync_lock = threading.Lock() self.auth_lock = threading.Lock() self.system_lock = threading.Lock() self._request_context = threading.local() self._system_credentials = self._load_system_credentials(environment_credentials) self.screener = ScreenerEngine(self.database) self.mentor_skills = MentorSkillRegistry(MENTOR_SKILLS_DIR) self.realtime_aggregator = WebRealtimeAggregator() self.screener.ensure_builtin_strategies() self._background_stop = threading.Event() self._background_thread = threading.Thread( target=self._background_refresh_loop, name="market-background-refresh", daemon=True, ) self._background_thread.start() def _load_system_credentials(self, environment: dict[str, str]) -> dict[str, Any]: encrypted = self.database.get_system_setting("credentials") current = self.vault.decrypt_json(encrypted) if encrypted else {} changed = False first_user_id = self.database.first_user_id() first_personal: dict[str, Any] = {} if first_user_id: first_encrypted = self.database.get_user_credentials(first_user_id) first_personal = self.vault.decrypt_json(first_encrypted) if first_encrypted else {} defaults = { "tushare_token": environment.get("tushare_token") or first_personal.get("tushare_token") or "", "platform_llm_primary_api_key": environment.get("platform_llm_primary_api_key") or first_personal.get("llm_primary_api_key") or "", "platform_llm_primary_base_url": environment.get("platform_llm_primary_base_url") or first_personal.get("llm_primary_base_url") or "https://api.openai.com/v1", "platform_llm_primary_model": environment.get("platform_llm_primary_model") or first_personal.get("llm_primary_model") or "", "platform_llm_fallback_api_key": environment.get("platform_llm_fallback_api_key") or first_personal.get("llm_fallback_api_key") or "", "platform_llm_fallback_base_url": environment.get("platform_llm_fallback_base_url") or first_personal.get("llm_fallback_base_url") or "", "platform_llm_fallback_model": environment.get("platform_llm_fallback_model") or first_personal.get("llm_fallback_model") or "", "member_daily_limit": 50, "background_refresh_enabled": True, } for key, value in defaults.items(): if key not in current: current[key] = value changed = True if not isinstance(current.get("llm_models"), list): migrated_models: list[dict[str, str]] = [] for role, label in (("primary", "原主模型"), ("fallback", "原辅助模型")): profile = { "api_key": str(current.get(f"platform_llm_{role}_api_key") or ""), "base_url": str(current.get(f"platform_llm_{role}_base_url") or ""), "model": str(current.get(f"platform_llm_{role}_model") or ""), } if profile["api_key"] or profile["model"]: model_id = f"migrated-{role}" migrated_models.append( {"id": model_id, "name": label, **profile} ) current[f"{role}_model_id"] = model_id current["llm_models"] = migrated_models current.setdefault("primary_model_id", "") current.setdefault("fallback_model_id", "") changed = True if changed or not encrypted: self.database.save_system_setting("credentials", self.vault.encrypt_json(current)) for row in self.database.list_user_credentials(): personal = self.vault.decrypt_json(str(row.get("encrypted_payload") or "")) if "tushare_token" in personal: personal.pop("tushare_token", None) self.database.save_user_credentials( int(row["user_id"]), self.vault.encrypt_json(personal) ) return current def _save_system_credentials(self, credentials: dict[str, Any]) -> None: with self.system_lock: self.database.save_system_setting("credentials", self.vault.encrypt_json(credentials)) self._system_credentials = dict(credentials) @property def configured(self) -> bool: return bool(self.token) def bind_user(self, user_id: int) -> None: self._request_context.user_id = int(user_id) encrypted = self.database.get_user_credentials(int(user_id)) self._request_context.credentials = self.vault.decrypt_json(encrypted) if encrypted else {} self._request_context.access = self.database.user_access(int(user_id)) or {} @property def current_user_id(self) -> int: user_id = getattr(self._request_context, "user_id", 0) if not user_id: raise ValueError("当前请求尚未绑定账号。") return int(user_id) def _credentials(self) -> dict[str, str]: credentials = getattr(self._request_context, "credentials", {}) return { "llm_primary_api_key": str(credentials.get("llm_primary_api_key") or ""), "llm_primary_base_url": str( credentials.get("llm_primary_base_url") or "https://api.openai.com/v1" ), "llm_primary_model": str(credentials.get("llm_primary_model") or ""), "llm_fallback_api_key": str(credentials.get("llm_fallback_api_key") or ""), "llm_fallback_base_url": str(credentials.get("llm_fallback_base_url") or ""), "llm_fallback_model": str(credentials.get("llm_fallback_model") or ""), } def _save_credentials(self, credentials: dict[str, str]) -> None: self.database.save_user_credentials( self.current_user_id, self.vault.encrypt_json(credentials), ) self._request_context.credentials = dict(credentials) @property def token(self) -> str: return str(self._system_credentials.get("tushare_token") or "") def _personal_llm_profile(self) -> dict[str, Any]: credentials = self._credentials() return { "source": "personal", "primary": { "api_key": credentials["llm_primary_api_key"], "base_url": credentials["llm_primary_base_url"], "model": credentials["llm_primary_model"], }, "fallback": { "api_key": credentials["llm_fallback_api_key"], "base_url": credentials["llm_fallback_base_url"], "model": credentials["llm_fallback_model"], }, } def _platform_llm_profile(self) -> dict[str, Any]: models = { str(item.get("id") or ""): item for item in self._system_credentials.get("llm_models") or [] if isinstance(item, dict) and item.get("id") } def selected(role: str) -> dict[str, str]: item = models.get(str(self._system_credentials.get(f"{role}_model_id") or ""), {}) return { "id": str(item.get("id") or ""), "name": str(item.get("name") or ""), "api_key": str(item.get("api_key") or ""), "base_url": str(item.get("base_url") or ""), "model": str(item.get("model") or ""), } return { "source": "platform", "primary": selected("primary"), "fallback": selected("fallback"), } @staticmethod def _profile_configured(profile: dict[str, str]) -> bool: return bool(profile.get("api_key") and profile.get("base_url") and profile.get("model")) def membership(self) -> dict[str, Any]: access = getattr(self._request_context, "access", {}) or self.database.user_access(self.current_user_id) or {} now = datetime.now(timezone.utc) starts = _parse_iso_datetime(access.get("membership_starts_at")) expires = _parse_iso_datetime(access.get("membership_expires_at")) subscribed = ( access.get("membership_status") == "active" and (not starts or starts <= now) and (not expires or expires > now) ) is_admin = str(access.get("role")) == "admin" active = is_admin or subscribed remaining_seconds = None if expires: remaining_seconds = max(0, int((expires - now).total_seconds())) return { "active": active, "subscribed": subscribed, "status": "active" if subscribed else str(access.get("membership_status") or "inactive"), "plan": str(access.get("membership_plan") or ""), "starts_at": str(access.get("membership_starts_at") or ""), "expires_at": str(access.get("membership_expires_at") or ""), "is_admin": is_admin, "remaining_seconds": remaining_seconds, "remaining_days": None if remaining_seconds is None else (remaining_seconds + 86399) // 86400, } def _resolved_llm_profile(self) -> dict[str, Any]: platform = self._platform_llm_profile() platform_ready = self.membership()["active"] and self._profile_configured(platform["primary"]) if platform_ready: return platform return {"source": "none", "primary": {}, "fallback": {}} @property def llm_primary_api_key(self) -> str: return str(self._resolved_llm_profile()["primary"].get("api_key") or "") @property def llm_primary_base_url(self) -> str: return str(self._resolved_llm_profile()["primary"].get("base_url") or "") @property def llm_primary_model(self) -> str: return str(self._resolved_llm_profile()["primary"].get("model") or "") @property def llm_fallback_api_key(self) -> str: return str(self._resolved_llm_profile()["fallback"].get("api_key") or "") @property def llm_fallback_base_url(self) -> str: return str(self._resolved_llm_profile()["fallback"].get("base_url") or "") @property def llm_fallback_model(self) -> str: return str(self._resolved_llm_profile()["fallback"].get("model") or "") @property def llm_source(self) -> str: return str(self._resolved_llm_profile().get("source") or "none") @property def llm_configured(self) -> bool: return bool(self.llm_primary_api_key and self.llm_primary_model) @property def llm_fallback_configured(self) -> bool: return bool( self.llm_fallback_api_key and self.llm_fallback_base_url and self.llm_fallback_model ) def save_llm_settings( self, primary: dict[str, Any], fallback: dict[str, Any], fallback_enabled: bool, ) -> None: personal = self._personal_llm_profile() primary_profile = self._validate_llm_profile( primary, personal["primary"], required=True, label="主模型", ) if fallback_enabled: fallback_profile = self._validate_llm_profile( fallback, personal["fallback"], required=True, label="辅助模型", ) else: fallback_profile = {"api_key": "", "base_url": "", "model": ""} credentials = self._credentials() credentials.update( { "llm_primary_api_key": primary_profile["api_key"], "llm_primary_base_url": primary_profile["base_url"], "llm_primary_model": primary_profile["model"], "llm_fallback_api_key": fallback_profile["api_key"], "llm_fallback_base_url": fallback_profile["base_url"], "llm_fallback_model": fallback_profile["model"], } ) self._save_credentials(credentials) def save_llm_mode(self, mode: str) -> None: raise ValueError("LLM 算力由管理员统一配置,会员账号自动使用平台模型。") def test_llm_profile(self, role: str, payload: dict[str, Any]) -> dict[str, Any]: personal = self._personal_llm_profile() if role == "primary": current = personal["primary"] label = "主模型" elif role == "fallback": current = personal["fallback"] label = "辅助模型" else: raise ValueError("模型角色不支持。") profile = self._validate_llm_profile(payload, current, required=True, label=label) try: return test_llm_connection(**profile) except LLMCompilerError as exc: raise ValueError(str(exc)) from exc @staticmethod def _validate_llm_profile( payload: dict[str, Any], current: dict[str, str], required: bool, label: str, ) -> dict[str, str]: api_key = str(payload.get("api_key") or current.get("api_key") or "").strip() base_url = str(payload.get("base_url") or current.get("base_url") or "").strip().rstrip("/") model = str(payload.get("model") or current.get("model") or "").strip() if not required and not any((api_key, base_url, model)): return {"api_key": "", "base_url": "", "model": ""} parsed = urlparse(base_url) if parsed.scheme not in {"http", "https"} or not parsed.netloc: raise ValueError(f"{label} Base URL 格式不正确。") if not api_key or len(api_key) > 300: raise ValueError(f"{label} API Key 不能为空或过长。") if not model or len(model) > 100: raise ValueError(f"{label}模型名称不能为空或过长。") return {"api_key": api_key, "base_url": base_url, "model": model} def llm_access_status(self) -> dict[str, Any]: platform = self._platform_llm_profile() membership = self.membership() limit = max(1, int(self._system_credentials.get("member_daily_limit") or 50)) used = self._platform_usage_today() if membership["active"] else 0 resolved = self._resolved_llm_profile() return { "mode": "platform" if membership["active"] else "locked", "resolved_source": resolved.get("source") or "none", "resolved_model": str(resolved.get("primary", {}).get("model") or ""), "platform_configured": self._profile_configured(platform["primary"]), "membership": membership, "daily_limit": limit, "used_today": used, "remaining_calls": None if membership["is_admin"] else max(0, limit - used), } def _platform_usage_today(self) -> int: now = datetime.now().astimezone() start = now.replace(hour=0, minute=0, second=0, microsecond=0).astimezone(timezone.utc) return self.database.count_llm_usage_since( self.current_user_id, "platform", start.isoformat(timespec="seconds"), ) def ensure_llm_access(self, feature: str) -> str: profile = self._resolved_llm_profile() source = str(profile.get("source") or "none") if source == "none" or not self._profile_configured(profile.get("primary") or {}): raise ValueError("请配置个人 LLM,或使用已开通会员的平台模型。") if source == "platform": limit = max(1, int(self._system_credentials.get("member_daily_limit") or 50)) if self._platform_usage_today() >= limit: raise ValueError(f"今日会员模型额度已用完({limit} 次)。") return source def record_llm_usage( self, feature: str, source: str, model: str, status: str, latency_ms: int = 0, ) -> None: self.database.record_llm_usage( self.current_user_id, feature, source, model, status, latency_ms ) def system_status(self) -> dict[str, Any]: platform = self._platform_llm_profile() model_pool = [] for item in self._system_credentials.get("llm_models") or []: if not isinstance(item, dict): continue profile = { "api_key": str(item.get("api_key") or ""), "base_url": str(item.get("base_url") or ""), "model": str(item.get("model") or ""), } model_pool.append( { "id": str(item.get("id") or ""), "name": str(item.get("name") or ""), "base_url": profile["base_url"], "model": profile["model"], "configured": self._profile_configured(profile), } ) return { "data": { "configured": self.configured, "background_refresh_enabled": bool( self._system_credentials.get("background_refresh_enabled", True) ), **self.database.status(), }, "llm": { "primary_configured": self._profile_configured(platform["primary"]), "fallback_configured": self._profile_configured(platform["fallback"]), "models": model_pool, "primary_model_id": str(self._system_credentials.get("primary_model_id") or ""), "fallback_model_id": str(self._system_credentials.get("fallback_model_id") or ""), }, "membership": { "member_daily_limit": max( 1, int(self._system_credentials.get("member_daily_limit") or 50) ) }, } def save_system_settings(self, payload: dict[str, Any]) -> dict[str, Any]: current = dict(self._system_credentials) token = str(payload.get("tushare_token") or current.get("tushare_token") or "").strip() if token and not TOKEN_PATTERN.fullmatch(token): raise ValueError("Tushare Token 格式不正确。") existing_models = { str(item.get("id") or ""): item for item in current.get("llm_models") or [] if isinstance(item, dict) and item.get("id") } raw_models = payload.get("models") models: list[dict[str, str]] = [] if raw_models is not None: if not isinstance(raw_models, list) or len(raw_models) > 20: raise ValueError("模型池格式不正确,最多可保存 20 个模型。") seen_ids: set[str] = set() seen_names: set[str] = set() for index, raw in enumerate(raw_models, start=1): if not isinstance(raw, dict): raise ValueError("模型池条目格式不正确。") model_id = str(raw.get("id") or f"model-{secrets.token_hex(6)}").strip() if not re.fullmatch(r"[A-Za-z0-9_-]{3,80}", model_id) or model_id in seen_ids: raise ValueError("模型 ID 不正确或重复。") name = validate_text(raw.get("name"), f"模型 {index} 名称", 50, required=True) normalized_name = name.casefold() if normalized_name in seen_names: raise ValueError("模型名称不能重复。") profile = self._validate_llm_profile( raw, existing_models.get(model_id) or {}, required=True, label=name, ) models.append({"id": model_id, "name": name, **profile}) seen_ids.add(model_id) seen_names.add(normalized_name) else: models = [dict(item) for item in existing_models.values()] model_ids = {item["id"] for item in models} primary_model_id = str( payload.get("primary_model_id", current.get("primary_model_id") or "") or "" ).strip() fallback_model_id = str( payload.get("fallback_model_id", current.get("fallback_model_id") or "") or "" ).strip() if models and primary_model_id not in model_ids: raise ValueError("请从模型池选择主模型。") if not models: primary_model_id = "" fallback_model_id = "" if fallback_model_id and fallback_model_id not in model_ids: raise ValueError("辅助模型不在模型池中。") if fallback_model_id and fallback_model_id == primary_model_id: raise ValueError("主模型与辅助模型不能相同。") try: daily_limit = max( 1, min( 1000, int(payload.get("member_daily_limit", current.get("member_daily_limit") or 50)), ), ) except (TypeError, ValueError) as exc: raise ValueError("会员每日额度应为 1 至 1000。") from exc current.update( { "tushare_token": token, "llm_models": models, "primary_model_id": primary_model_id, "fallback_model_id": fallback_model_id, "member_daily_limit": daily_limit, "background_refresh_enabled": bool( payload.get( "background_refresh_enabled", current.get("background_refresh_enabled", True), ) ), } ) self._save_system_credentials(current) return self.system_status() def test_system_llm_profile(self, model_id: str, payload: dict[str, Any]) -> dict[str, Any]: current = next( ( item for item in self._system_credentials.get("llm_models") or [] if str(item.get("id") or "") == model_id ), {}, ) label = validate_text(payload.get("name") or current.get("name"), "模型名称", 50, required=True) profile = self._validate_llm_profile( payload, current, required=True, label=label ) try: return test_llm_connection(**profile) except LLMCompilerError as exc: raise ValueError(str(exc)) from exc def admin_users(self) -> list[dict[str, Any]]: original_user_id = getattr(self._request_context, "user_id", 0) original_credentials = getattr(self._request_context, "credentials", {}) original_access = getattr(self._request_context, "access", {}) rows = [] try: for user in self.database.list_users(): self._request_context.user_id = int(user["id"]) self._request_context.access = user membership = self.membership() used = self._platform_usage_today() if membership["active"] else 0 rows.append({ **user, "membership_active": membership["active"], "membership_subscribed": membership["subscribed"], "used_today": used, }) finally: self._request_context.user_id = original_user_id self._request_context.credentials = original_credentials self._request_context.access = original_access return rows def update_membership(self, payload: dict[str, Any]) -> None: try: user_id = int(payload.get("user_id")) except (TypeError, ValueError) as exc: raise ValueError("会员账号不正确。") from exc status = str(payload.get("status") or "inactive") if status not in {"active", "inactive", "suspended"}: raise ValueError("会员状态不正确。") access = self.database.user_access(user_id) if not access: raise ValueError("用户不存在。") starts_at = None expires_at = None plan = "" if status == "active": duration = str(payload.get("duration") or "").strip() durations = { "1_month": (1, "1个月"), "3_months": (3, "3个月"), "12_months": (12, "12个月"), "3_years": (36, "3年"), "permanent": (0, "永久"), } if duration not in durations: raise ValueError("请选择会员开通时长。") now = datetime.now(timezone.utc) existing_start = _parse_iso_datetime(access.get("membership_starts_at")) existing_expiry = _parse_iso_datetime(access.get("membership_expires_at")) starts = existing_start if existing_start and existing_start <= now else now months, plan = durations[duration] starts_at = starts.isoformat(timespec="seconds") if months: renewal_base = existing_expiry if existing_expiry and existing_expiry > now else now expires_at = _add_months(renewal_base, months).isoformat(timespec="seconds") if not self.database.update_membership( user_id, status, plan, starts_at, expires_at ): raise ValueError("用户不存在。") def request_background_sync(self, trade_date: str) -> bool: if self.sync_lock.locked(): return False normalized = normalize_date(trade_date) threading.Thread( target=self._run_background_sync, args=(normalized,), name=f"market-sync-{normalized}", daemon=True, ).start() return True def _run_background_sync(self, trade_date: str) -> None: try: self.sync_dashboard(trade_date) except Exception: return def _background_refresh_loop(self) -> None: self._background_stop.wait(3) while not self._background_stop.is_set(): try: if self.configured and self._system_credentials.get("background_refresh_enabled", True): today = date.today().strftime("%Y%m%d") snapshot = self.database.get_snapshot(today) or {} if self._realtime_snapshot_due(today, snapshot): self._run_background_sync(today) except Exception: pass self._background_stop.wait(5) def register_account(self, username: str, password: str) -> dict[str, Any]: username = username.strip() self._validate_account_input(username, password) with self.auth_lock: salt, password_digest = hash_password(password) user = self.database.create_user(username, salt, password_digest) return self.create_account_session(user) def login_account(self, username: str, password: str) -> dict[str, Any]: username = username.strip() if not username or not password: raise ValueError("账号名和密码不能为空。") user = self.database.user_by_username(username) if not user or not verify_password( password, str(user.get("password_salt") or ""), str(user.get("password_hash") or ""), ): raise ValueError("账号名或密码不正确。") return self.create_account_session(user) def change_password(self, current_password: str, new_password: str) -> None: current_password = str(current_password or "") self._validate_account_input(str(self.database.user_access(self.current_user_id)["username"]), new_password) credentials = self.database.user_password(self.current_user_id) if not credentials or not verify_password( current_password, str(credentials.get("password_salt") or ""), str(credentials.get("password_hash") or ""), ): raise ValueError("当前密码不正确。") salt, digest = hash_password(new_password) if not self.database.update_user_password(self.current_user_id, salt, digest): raise ValueError("账号不存在。") def create_account_session(self, user: dict[str, Any]) -> dict[str, Any]: session_token = secrets.token_urlsafe(32) csrf_token = secrets.token_urlsafe(24) expires = datetime.now(timezone.utc) + timedelta(seconds=SESSION_MAX_AGE) self.database.create_session( token_hash(session_token), int(user["id"]), csrf_token, expires.isoformat(timespec="seconds"), ) self.bind_user(int(user["id"])) access = self.database.user_access(int(user["id"])) or {} return { "user": { "id": int(user["id"]), "username": str(user["username"]), "role": str(access.get("role") or "user"), "membership": self.membership(), }, "session_token": session_token, "csrf_token": csrf_token, } @staticmethod def _validate_account_input(username: str, password: str) -> None: if not USERNAME_PATTERN.fullmatch(username): raise ValueError("账号名应为 3 至 30 位中文、字母、数字、下划线或连字符。") if len(password) < 8 or len(password) > 128: raise ValueError("密码长度应为 8 至 128 位。") if password.isalpha() or password.isdigit(): raise ValueError("密码应同时包含字母、数字或符号中的至少两类。") def save_birth_profile(self, payload: dict[str, Any]) -> dict[str, Any]: birth_datetime = str(payload.get("birth_datetime") or "").strip() gender = str(payload.get("gender") or "unspecified").strip() current_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) personal = build_personal_field(birth_datetime, gender, current_date) encrypted = self.vault.encrypt_json( {"birth_datetime": birth_datetime, "gender": gender} ) self.database.save_user_birth_profile(self.current_user_id, encrypted) return self._public_personal_profile(personal) def stored_birth_profile(self) -> dict[str, str] | None: encrypted = self.database.get_user_birth_profile(self.current_user_id) if not encrypted: return None payload = self.vault.decrypt_json(encrypted) birth_datetime = str(payload.get("birth_datetime") or "").strip() if not birth_datetime: return None return { "birth_datetime": birth_datetime, "gender": str(payload.get("gender") or "unspecified"), } def account_personal_field( self, current_date: str, current_field: dict[str, Any], public: bool = False, ) -> dict[str, Any] | None: stored = self.stored_birth_profile() if not stored: return None personal = build_personal_field( stored["birth_datetime"], stored["gender"], current_date, current_field, ) if public: return self._public_personal_profile(personal) personal.pop("birth", None) return personal @staticmethod def _public_personal_profile(personal: dict[str, Any]) -> dict[str, Any]: allowed = { "day_master", "ten_god_tendency", "element_balance", "balance_tendency", "current", "notice", } return {key: value for key, value in personal.items() if key in allowed} def get_dashboard(self, trade_date: str, force: bool = False) -> dict[str, Any]: normalized_date = normalize_date(trade_date) now = datetime.now().astimezone() if ( normalized_date == now.strftime("%Y%m%d") and now.time().replace(tzinfo=None) < datetime.strptime("09:15", "%H:%M").time() ): previous = self.database.get_latest_real_snapshot(normalized_date, strictly_before=True) if previous: carried = self._carry_dashboard(previous, normalized_date, "盘前沿用最近交易日收盘行情") return self._apply_reason_overrides(self._with_storage(carried, cached=True)) if not force: snapshot = self.database.get_snapshot(normalized_date) if snapshot and str((snapshot.get("meta") or {}).get("source") or "") != "demo": snapshot = copy.deepcopy(snapshot) if normalized_date != now.strftime("%Y%m%d"): snapshot.setdefault("meta", {}).update( {"realtime": False, "market_status": "closed"} ) snapshot = self._enrich_dashboard_sentiment(snapshot, normalized_date) snapshot.setdefault("meta", {})["requested_date"] = self._display_compact_date(normalized_date) return self._apply_reason_overrides(self._with_storage(snapshot, cached=True)) return self.sync_dashboard(normalized_date) @staticmethod def _display_compact_date(compact: str) -> str: return f"{compact[:4]}-{compact[4:6]}-{compact[6:8]}" def _carry_dashboard( self, snapshot: dict[str, Any], requested_date: str, reason: str ) -> dict[str, Any]: carried = copy.deepcopy(snapshot) meta = carried.setdefault("meta", {}) meta.update( { "requested_date": self._display_compact_date(requested_date), "carried_forward": True, "realtime": False, "market_status": "closed", "notice": reason, } ) return carried def _realtime_snapshot_due( self, normalized_date: str, snapshot: dict[str, Any], ) -> bool: if not self.configured or normalized_date != date.today().strftime("%Y%m%d"): return False now = datetime.now().astimezone() local_time = now.time().replace(tzinfo=None) realtime_start = datetime.strptime("09:15", "%H:%M").time() morning_end = datetime.strptime("11:35", "%H:%M").time() afternoon_start = datetime.strptime("12:55", "%H:%M").time() realtime_end = datetime.strptime("15:05", "%H:%M").time() in_session = ( realtime_start <= local_time < morning_end or afternoon_start <= local_time < realtime_end ) if not in_session: return False meta = snapshot.get("meta") or {} snapshot_trade_date = str(meta.get("trade_date") or "").replace("-", "") if snapshot_trade_date and snapshot_trade_date != normalized_date: return False if not meta.get("realtime"): return True try: updated_at = datetime.fromisoformat(str(meta.get("updated_at") or "")) if updated_at.tzinfo is None: updated_at = updated_at.replace(tzinfo=now.tzinfo) except ValueError: return True age_seconds = (now - updated_at.astimezone(now.tzinfo)).total_seconds() return age_seconds >= 8 def sync_dashboard(self, trade_date: str) -> dict[str, Any]: normalized_date = normalize_date(trade_date) source = "tushare" with self.sync_lock: sync_id = self.database.start_sync(normalized_date, source) try: if not self.configured: raise TushareError("公共行情尚未配置") dashboard = TushareClient(self.token).dashboard(normalized_date) dashboard["meta"]["source"] = source dashboard["meta"]["requested_date"] = self._display_compact_date(normalized_date) dashboard = self._enrich_dashboard_sentiment(dashboard, normalized_date) record_count = self._record_count(dashboard) actual_date = normalize_date( str(dashboard.get("meta", {}).get("trade_date") or normalized_date) ) self.database.save_snapshot(actual_date, source, dashboard) self.database.finish_sync( sync_id, "success", record_count, dashboard.get("meta", {}).get("notice", ""), source, ) return self._apply_reason_overrides(self._with_storage(dashboard, cached=False)) except TushareError as exc: fallback = self.database.get_latest_real_snapshot(normalized_date) if fallback: carried = self._carry_dashboard( fallback, normalized_date, f"最新行情暂不可用,沿用最近收盘快照:{exc}" ) self.database.finish_sync( sync_id, "fallback", self._record_count(carried), str(exc), "tushare" ) return self._apply_reason_overrides(self._with_storage(carried, cached=True)) self.database.finish_sync(sync_id, "failed", message=str(exc)) raise ValueError("暂无可用的真实行情快照,请等待后台完成首次同步。") from exc except Exception as exc: self.database.finish_sync(sync_id, "failed", message=str(exc)) raise def _enrich_dashboard_sentiment( self, dashboard: dict[str, Any], end_date: str, ) -> dict[str, Any]: history = self.database.list_snapshot_payloads(end_date, 240) return apply_sentiment_to_dashboard(dashboard, history) def sentiment_history(self, trade_date: str, limit: int = 20) -> dict[str, Any]: normalized_date = normalize_date(trade_date) limit = max(10, min(120, int(limit))) full_series = build_sentiment_history( self.database.list_snapshot_payloads(normalized_date, 240) ) series = latest_contiguous_history(full_series) rows = series[-limit:] return { "trade_date": rows[-1]["trade_date"] if rows else normalized_date, "available_days": len(series), "stored_days": len(full_series), "requested_days": limit, "rows": rows, "weights": COMPONENT_WEIGHTS, "normalization": rows[-1]["normalization"] if rows else "固定锚点", } def rotation_history(self, trade_date: str, limit: int = 9) -> dict[str, Any]: normalized_date = normalize_date(trade_date) # 板块轮动固定展示最近 9 个交易日,按由近到远排列。 limit = 9 snapshots = self.database.list_snapshot_payloads(normalized_date, 240) by_trade_date: dict[str, dict[str, Any]] = {} for snapshot in snapshots: meta = snapshot.get("meta") or {} actual_date = str(meta.get("trade_date") or snapshot.get("_snapshot_date") or "") compact_date = actual_date.replace("-", "") if len(compact_date) == 8: by_trade_date[compact_date] = snapshot sentiment_dates = { str(row.get("trade_date") or "").replace("-", "") for row in latest_contiguous_history(build_sentiment_history(snapshots)) } ordered_dates = sorted( date_key for date_key in by_trade_date if not sentiment_dates or date_key in sentiment_dates )[-limit:][::-1] rows = [] for date_key in ordered_dates: snapshot = by_trade_date[date_key] sector_context = { str(item.get("name") or ""): item for item in snapshot.get("sectors") or [] } sectors = [] for item in (snapshot.get("sector_rotation") or [])[:12]: name = str(item.get("name") or "").strip() context = sector_context.get(name, {}) sectors.append( { "name": name, "rank": int(item.get("rank") or len(sectors) + 1), "trend": item.get("trend") or "持平", "count": int(item.get("count") or 0), "strength": float(item.get("strength") or context.get("strength") or 0), "change": float(context.get("change") or 0), "leader": item.get("leader") or context.get("leader") or "--", } ) rows.append( { "trade_date": f"{date_key[:4]}-{date_key[4:6]}-{date_key[6:]}", "sectors": sectors, } ) return { "trade_date": rows[0]["trade_date"] if rows else normalized_date, "available_days": len(ordered_dates), "requested_days": limit, "rows": rows, } def status(self) -> dict[str, Any]: llm_access = self.llm_access_status() return { "configured": self.configured, "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, "llm_fallback_model": self.llm_fallback_model if self.llm_fallback_configured else "", "llm_access": llm_access, "birth_profile_configured": bool(self.stored_birth_profile()), "birth_profile": self.stored_birth_profile(), **self.database.status(), } def realtime_aggregate_health(self, sector: str = "") -> dict[str, Any]: sector = validate_text(sector, "板块名称", 50) return self.realtime_aggregator.health_snapshot(sector) def screener_setup(self, trade_date: str) -> dict[str, Any]: normalized_date = normalize_date(trade_date) regime = self.screener.detect_regime(normalized_date) factor_dates = self.database.factor_dates(normalized_date, 100) return { "trade_date": normalized_date, "regime": regime, "regimes": [{"id": key, "label": value} for key, value in REGIMES.items()], "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), "start_date": factor_dates[0] if factor_dates else "", "end_date": factor_dates[-1] if factor_dates else "", "ready": len(factor_dates) >= 21, }, "llm": { "configured": self.llm_configured, "model": self.llm_primary_model if self.llm_configured else "", "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(self.current_user_id, normalized_date), } def sync_screener_data(self, trade_date: str, lookback: int = 45) -> dict[str, Any]: if not self.configured: raise ValueError("请先配置 Tushare Token。") normalized_date = normalize_date(trade_date) lookback = max(25, min(80, int(lookback))) with self.sync_lock: return FactorDataService(self.database, TushareClient(self.token)).sync( normalized_date, lookback ) def compile_screener_strategy(self, prompt: str, regime: str) -> dict[str, Any]: prompt = prompt.strip() if not prompt or len(prompt) > 3000: raise ValueError("策略描述应为 1 至 3000 个字符。") if regime not in REGIMES: raise ValueError("市场阶段不支持。") notice = "" compiled = None primary_error = "" source = self.llm_source started = datetime.now(timezone.utc) if source == "platform": self.ensure_llm_access("screener") if self.llm_configured: try: compiled = compile_strategy_with_llm( prompt, regime, self.llm_primary_api_key, self.llm_primary_base_url, self.llm_primary_model, ) except LLMCompilerError as exc: primary_error = str(exc) if compiled is None and self.llm_fallback_configured: try: compiled = compile_strategy_with_llm( prompt, regime, self.llm_fallback_api_key, self.llm_fallback_base_url, self.llm_fallback_model, ) compiled["compiler"] = "llm_fallback" notice = f"主模型调用失败,已自动切换辅助模型。{primary_error}" if primary_error else "已使用辅助模型编译。" except LLMCompilerError as exc: fallback_error = str(exc) compiled = compile_local_strategy(prompt, regime) notice = f"主模型和辅助模型均不可用,已使用本地模板。主模型:{primary_error or '未配置'};辅助模型:{fallback_error}" if compiled is None: compiled = compile_local_strategy(prompt, regime) notice = ( f"主模型调用失败且未配置辅助模型,已使用本地模板:{primary_error}" if primary_error else "尚未配置 LLM,当前使用本地受控模板编译。" ) compiled["formula"] = self.screener.validate_formula(compiled["formula"]) compiled["notice"] = notice if source in {"personal", "platform"}: elapsed = int((datetime.now(timezone.utc) - started).total_seconds() * 1000) status = "success" if str(compiled.get("compiler") or "").startswith("llm") else "failed" self.record_llm_usage( "screener", source, str(compiled.get("model") or self.llm_primary_model), status, elapsed, ) return compiled def save_screener_strategy(self, payload: dict[str, Any]) -> dict[str, Any]: name = validate_text(payload.get("name"), "策略名称", 60, required=True) description = validate_text(payload.get("description"), "策略说明", 1000) regimes = payload.get("regimes") or [] 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( 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(self.current_user_id, strategy_id) return { "deleted": deleted, "strategies": self.database.list_screener_strategies(self.current_user_id), } def mentor_setup(self, trade_date: str) -> dict[str, Any]: normalized_date = normalize_date(trade_date) mentors = [skill.public() for skill in self.mentor_skills.list_skills()] if not mentors: raise ValueError("游资skills 目录中没有可用的 SKILL.md。") snapshot = self.database.get_snapshot(normalized_date) actual_date = str((snapshot or {}).get("meta", {}).get("trade_date") or normalized_date) return { "trade_date": actual_date, "mentors": mentors, "llm": { "configured": self.llm_configured, "model": self.llm_primary_model if self.llm_configured else "", "fallback_configured": self.llm_fallback_configured, "fallback_model": self.llm_fallback_model if self.llm_fallback_configured else "", }, } def mentor_chat(self, payload: dict[str, Any]) -> dict[str, Any]: mentor_id = validate_text(payload.get("mentor_id"), "问师角色", 100, required=True) question = validate_text(payload.get("question"), "问题", 2000, required=True) trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) history = self._validate_mentor_history(payload.get("history") or []) skill = self.mentor_skills.get_skill(mentor_id) context = self._build_mentor_context(trade_date, question) source = self.ensure_llm_access("mentor") primary_error = "" result = None compiler = "primary" if self.llm_configured: try: result = chat_with_mentor( skill, context, question, history, self.llm_primary_api_key, self.llm_primary_base_url, self.llm_primary_model, ) except MentorAgentError as exc: primary_error = str(exc) if result is None and self.llm_fallback_configured: try: result = chat_with_mentor( skill, context, question, history, self.llm_fallback_api_key, self.llm_fallback_base_url, self.llm_fallback_model, ) compiler = "fallback" except MentorAgentError as exc: fallback_error = str(exc) self.record_llm_usage( "mentor", source, self.llm_fallback_model, "failed" ) raise ValueError( f"主模型和辅助模型均不可用。主模型:{primary_error or '未配置'};" f"辅助模型:{fallback_error}" ) from exc if result is None: self.record_llm_usage("mentor", source, self.llm_primary_model, "failed") raise ValueError(f"主模型不可用且未配置辅助模型:{primary_error}") self.record_llm_usage( "mentor", source, str(result.get("model") or ""), "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(), "compiler": compiler, "requested_trade_date": trade_date, "data_trade_date": context["data_trade_date"], "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" fields = { "stock_amount_percentile": {"line": 1, "label": "成交额全市场分位", "unit": "%", "min": 0, "max": 100}, "stock_turnover_rate": {"line": 1, "label": "个股换手率", "unit": "%", "min": 0, "max": 100}, "stock_turnover_relative": {"line": 1, "label": "相对市场换手", "unit": "倍", "min": 0, "max": 20}, "stock_volume_activity_ratio": {"line": 1, "label": "同进度量能", "unit": "倍", "min": 0, "max": 20}, "stock_seal_amount_million": {"line": 1, "label": "封单金额", "unit": "万元", "min": 0, "max": 100000000}, "stock_open_times": {"line": 1, "label": "开板次数", "unit": "次", "min": 0, "max": 100, "integer": True}, "stock_change": {"line": 2, "label": "个股涨跌幅", "unit": "%", "min": -100, "max": 100}, "stock_streak": {"line": 2, "label": "连板高度", "unit": "板", "min": 0, "max": 100, "integer": True}, "stock_status": {"line": 2, "label": "个股状态", "type": "select", "options": ["普通", "涨停", "炸板", "跌停"]}, "sector_name": {"line": [3, 4], "label": "申万二级行业", "type": "text", "max_length": 50}, "sector_up_count": {"line": 3, "label": "行业上涨家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "sector_down_count": {"line": 3, "label": "行业下跌家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "sector_coverage": {"line": 3, "label": "成分行情覆盖率", "unit": "%", "min": 0, "max": 100}, "sector_relative_turnover": {"line": 3, "label": "行业相对市场换手", "unit": "倍", "min": 0, "max": 20}, "sector_member_equal_change": {"line": 3, "label": "成分等权涨跌幅", "unit": "%", "min": -100, "max": 100}, "sector_change": {"line": 4, "label": "申万官方涨跌幅", "unit": "%", "min": -100, "max": 100}, "sector_leading_pct": {"line": [3, 4], "label": "行业领涨股涨跌幅", "unit": "%", "min": -100, "max": 100}, "market_sentiment_score": {"line": 5, "label": "市场情绪温度", "unit": "分", "min": 0, "max": 100}, "market_seal_rate": {"line": 5, "label": "封板率", "unit": "%", "min": 0, "max": 100}, "market_amount_billion": {"line": 5, "label": "两市成交额", "unit": "亿元", "min": 0, "max": 10000000}, "market_recent_average_amount_billion": {"line": 5, "label": "近期平均成交额", "unit": "亿元", "min": 0, "max": 10000000}, "market_up_count": {"line": 5, "label": "上涨家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "market_down_count": {"line": 5, "label": "下跌家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "market_limit_up_count": {"line": 5, "label": "涨停家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "market_limit_down_count": {"line": 5, "label": "跌停家数", "unit": "家", "min": 0, "max": 10000, "integer": True}, "index_sh_change": {"line": 6, "label": "上证指数涨跌幅", "unit": "%", "min": -20, "max": 20}, "index_sz_change": {"line": 6, "label": "深证成指涨跌幅", "unit": "%", "min": -20, "max": 20}, "index_cy_change": {"line": 6, "label": "创业板指涨跌幅", "unit": "%", "min": -20, "max": 20}, "note": {"line": [], "label": "补录说明", "type": "text", "max_length": 200}, } if intraday: for key in ("stock_seal_amount_million", "stock_open_times"): fields.pop(key) else: for key in ("stock_turnover_relative", "stock_volume_activity_ratio", "sector_relative_turnover"): fields.pop(key) return fields @classmethod def _validate_heaven_manual_data( cls, raw: Any, market_mode: str ) -> dict[str, Any]: if raw in (None, ""): return {} if not isinstance(raw, dict): raise ValueError("六爻补录数据格式不正确。") schema = cls._heaven_manual_schema(market_mode) unknown = set(raw) - set(schema) if unknown: raise ValueError(f"六爻补录包含未知字段:{next(iter(sorted(unknown)))}") values: dict[str, Any] = {} for key, value in raw.items(): if value is None or (isinstance(value, str) and not value.strip()): continue spec = schema[key] if spec.get("type") == "text": values[key] = validate_text(value, spec["label"], int(spec["max_length"])) continue if spec.get("type") == "select": text = str(value).strip() if text not in spec["options"]: raise ValueError(f"{spec['label']}不在允许范围内。") values[key] = text continue try: number = float(value) except (TypeError, ValueError) as exc: raise ValueError(f"{spec['label']}必须是数字。") from exc if number < float(spec["min"]) or number > float(spec["max"]): raise ValueError( f"{spec['label']}应在 {spec['min']} 至 {spec['max']} 之间。" ) values[key] = int(number) if spec.get("integer") else number return values @staticmethod def _apply_heaven_manual_data( dashboard: dict[str, Any], index_context: dict[str, Any], sector: dict[str, Any] | None, stock: dict[str, Any] | None, manual_data: dict[str, Any], market_mode: str, trade_date: str, stock_code: str, ) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any], dict[str, Any]]: dashboard = copy.deepcopy(dashboard) index_context = copy.deepcopy(index_context or {}) sector = copy.deepcopy(sector or {}) stock = copy.deepcopy(stock or {}) overview = dashboard.setdefault("overview", {}) stock_map = { "stock_amount_percentile": "amount_percentile", "stock_turnover_rate": "turnover_rate", "stock_turnover_relative": "turnover_relative", "stock_volume_activity_ratio": "volume_activity_ratio", "stock_seal_amount_million": "seal_amount_million", "stock_open_times": "open_times", "stock_change": "change", "stock_streak": "streak", "stock_status": "status", } sector_map = { "sector_name": "name", "sector_up_count": "up_count", "sector_down_count": "down_count", "sector_coverage": "coverage", "sector_relative_turnover": "relative_turnover", "sector_member_equal_change": "member_equal_change", "sector_change": "change", "sector_leading_pct": "leading_pct", } overview_map = { "market_sentiment_score": "sentiment_score", "market_seal_rate": "seal_rate", "market_amount_billion": "amount_billion", "market_recent_average_amount_billion": "recent_average_amount_billion", "market_up_count": "up_count", "market_down_count": "down_count", "market_limit_up_count": "limit_up_count", "market_limit_down_count": "limit_down_count", } for manual_key, target in stock_map.items(): if manual_key in manual_data: stock[target] = manual_data[manual_key] for manual_key, target in sector_map.items(): if manual_key in manual_data: sector[target] = manual_data[manual_key] for manual_key, target in overview_map.items(): if manual_key in manual_data: overview[target] = manual_data[manual_key] if any(key.startswith("stock_") for key in manual_data): stock.setdefault("code", stock_code) stock.setdefault("name", stock_code or "--") stock["_quantitative_mode"] = "intraday" if market_mode == "intraday" else "historical" if market_mode == "intraday" and "stock_volume_activity_ratio" in manual_data: stock["activity_source"] = "user_supplied" if any(key.startswith("sector_") for key in manual_data): sector["_quantitative_mode"] = "intraday" if market_mode == "intraday" else "historical" sector.setdefault("taxonomy", "sw_l2") index_keys = ( ("index_sh_change", "000001.SH", "上证指数"), ("index_sz_change", "399001.SZ", "深证成指"), ("index_cy_change", "399006.SZ", "创业板指"), ) rows = {str(row.get("ts_code") or row.get("code") or ""): dict(row) for row in index_context.get("indices") or []} for manual_key, code, name in index_keys: if manual_key not in manual_data: continue row = rows.get(code, {"ts_code": code, "name": name}) row.update({"pct_chg": manual_data[manual_key], "trade_date": trade_date}) rows[code] = row ordered_rows = [rows.get(code) for _, code, _ in index_keys] if all(ordered_rows): index_context["indices"] = ordered_rows changes = [float(row.get("pct_chg") or 0) for row in ordered_rows] aggregate = dict(index_context.get("aggregate") or {}) aggregate["average_pct_chg"] = sum(changes) / 3 index_context["aggregate"] = aggregate return dashboard, index_context, sector, stock @classmethod def _heaven_line_checks( cls, trade_date: str, dashboard: dict[str, Any], recent_history: list[dict[str, Any]], index_context: dict[str, Any], sector: dict[str, Any], stock: dict[str, Any], market_mode: str, manual_data: dict[str, Any], ) -> list[dict[str, Any]]: intraday = market_mode == "intraday" closed = market_mode == "closed" schema = cls._heaven_manual_schema(market_mode) required = { 1: (["stock_amount_percentile", "stock_turnover_relative", "stock_volume_activity_ratio"] if intraday else ["stock_amount_percentile", "stock_turnover_rate", "stock_seal_amount_million", "stock_open_times"]), 2: ["stock_change", "stock_streak", "stock_status"], 3: (["sector_name", "sector_up_count", "sector_down_count", "sector_coverage", "sector_relative_turnover"] if intraday else ["sector_name", "sector_up_count", "sector_down_count", "sector_coverage", "sector_member_equal_change", "sector_leading_pct"]), 4: ["sector_name", "sector_change", "sector_leading_pct"], 5: ["market_sentiment_score", "market_seal_rate", "market_amount_billion", "market_recent_average_amount_billion", "market_up_count", "market_down_count", "market_limit_up_count", "market_limit_down_count"], 6: ["index_sh_change", "index_sz_change", "index_cy_change"], } names = { 1: ("初爻", "个股内核", "成交活跃、换手与量能"), 2: ("二爻", "个股外显", "涨跌、连板与状态"), 3: ("三爻", "行业内核", "行业宽度与成交活跃"), 4: ("四爻", "行业外显", "行业涨跌与领涨表现"), 5: ("五爻", "市场内核", "情绪、封板、成交与市场宽度"), 6: ("上爻", "指数外显", "三大指数当日涨跌"), } index_date = str(index_context.get("trade_date") or "").replace("-", "") index_rows = list(index_context.get("indices") or []) index_dates = {str(row.get("trade_date") or "").replace("-", "") for row in index_rows} index_issues = [] if len(index_rows) < 3: index_issues.append(f"三大指数仅取得 {len(index_rows)}/3 条行情") elif index_date != trade_date or index_dates != {trade_date}: actual_dates = "、".join(sorted(value for value in index_dates if value)) or "未知" index_issues.append(f"指数实际日期为 {actual_dates},目标交易日为 {trade_date}") elif not index_context.get("precise"): index_issues.append("三大指数行情未通过完整性校验") elif intraday and not index_context.get("realtime"): index_issues.append("盘中缺少可核验的实时指数行情") elif not intraday and (index_context.get("realtime") or str(index_context.get("source") or "") != "tushare"): index_issues.append("收盘或历史行情不是官方指数日线") sector_date = str(sector.get("trade_date") or "").replace("-", "") sector_coverage = float(sector.get("coverage") or 0) sector_common = [] if not sector: sector_common.append("未取得申万二级行业归属") elif sector.get("taxonomy") != "sw_l2": sector_common.append("行业分类不是申万二级") elif sector_date != trade_date: sector_common.append("行业行情日期与目标交易日不一致") elif intraday and not sector.get("realtime"): sector_common.append("盘中行业行情不是申万实时行情") elif market_mode == "historical" and sector.get("realtime"): sector_common.append("历史行业行情不能使用实时快照") elif closed and sector.get("realtime") and not sector.get("finalized"): sector_common.append("收盘行业实时行情尚未形成15:00最终快照") sector_inner = list(sector_common) sector_outer = list(sector_common) if not sector.get("inner_precise", sector.get("precise")): sector_inner.append(str(sector.get("inner_error") or sector.get("error") or "行业内核数据未通过校验")) if not sector.get("outer_precise", sector.get("precise")): sector_outer.append(str(sector.get("outer_error") or sector.get("error") or "行业外显数据未通过校验")) if sector and sector_coverage < 90: sector_inner.append(f"行业成分行情覆盖率仅 {sector_coverage:.1f}%,低于 90%") if sector.get("realtime") and not sector.get("relative_turnover"): sector_inner.append("缺少行业相对全市场换手活跃度") stock_date = str(stock.get("trade_date") or "").replace("-", "") stock_common = [] if not stock.get("code"): stock_common.append("尚未载入有效个股") elif stock_date != trade_date: stock_common.append(f"个股实际日期为 {stock_date or '未知'},目标交易日为 {trade_date}") elif not stock.get("precise"): stock_common.append("个股行情未通过完整性校验") elif intraday and not stock.get("realtime"): stock_common.append("盘中个股行情不是实时行情") elif not intraday and (stock.get("realtime") or str(stock.get("data_source") or "") != "tushare"): stock_common.append("收盘或历史个股行情不是官方日线") stock_inner = list(stock_common) if intraday and stock.get("turnover_source") in {None, "", "unavailable"}: stock_inner.append("缺少可核验的实时换手率") if intraday and stock.get("activity_source") in {None, "", "unavailable"}: stock_inner.append("缺少同时间进度量能基准") overview = dashboard.get("overview") or {} market_key_map = { "market_sentiment_score": "sentiment_score", "market_seal_rate": "seal_rate", "market_amount_billion": "amount_billion", "market_recent_average_amount_billion": "recent_average_amount_billion", "market_up_count": "up_count", "market_down_count": "down_count", "market_limit_up_count": "limit_up_count", "market_limit_down_count": "limit_down_count", } market_issues = [] for manual_key, source_key in market_key_map.items(): if source_key == "recent_average_amount_billion": history_values = [item.get("amount_billion") for item in recent_history[:-1] if item.get("amount_billion") is not None] if source_key not in overview and not history_values: market_issues.append(f"缺少{schema[manual_key]['label']}") elif source_key not in overview or overview.get(source_key) is None: market_issues.append(f"缺少{schema[manual_key]['label']}") automatic_issues = { 1: stock_inner, 2: stock_common, 3: sector_inner, 4: sector_outer, 5: market_issues, 6: index_issues, } limits = list(dashboard.get("limits") or []) scores = _market_line_scores(dashboard, recent_history, index_context, sector, stock, limits) value_map: dict[str, Any] = { "stock_amount_percentile": stock.get("amount_percentile"), "stock_turnover_rate": stock.get("turnover_rate"), "stock_turnover_relative": stock.get("turnover_relative"), "stock_volume_activity_ratio": stock.get("volume_activity_ratio"), "stock_seal_amount_million": stock.get("seal_amount_million"), "stock_open_times": stock.get("open_times"), "stock_change": stock.get("change"), "stock_streak": stock.get("streak"), "stock_status": stock.get("status"), "sector_name": sector.get("name"), "sector_up_count": sector.get("up_count"), "sector_down_count": sector.get("down_count"), "sector_coverage": sector.get("coverage"), "sector_relative_turnover": sector.get("relative_turnover"), "sector_member_equal_change": sector.get("member_equal_change"), "sector_change": sector.get("change"), "sector_leading_pct": sector.get("leading_pct"), "market_sentiment_score": overview.get("sentiment_score"), "market_seal_rate": overview.get("seal_rate"), "market_amount_billion": overview.get("amount_billion"), "market_recent_average_amount_billion": overview.get("recent_average_amount_billion"), "market_up_count": overview.get("up_count"), "market_down_count": overview.get("down_count"), "market_limit_up_count": overview.get("limit_up_count"), "market_limit_down_count": overview.get("limit_down_count"), } history_values = [float(item.get("amount_billion")) for item in recent_history[:-1] if item.get("amount_billion") is not None] if value_map["market_recent_average_amount_billion"] is None and history_values: value_map["market_recent_average_amount_billion"] = sum(history_values) / len(history_values) if value_map["stock_amount_percentile"] is None and not intraday: amount = float(stock.get("amount_billion") or 0) amounts = [float(item.get("amount_billion") or 0) for item in limits if item.get("amount_billion") is not None] value_map["stock_amount_percentile"] = ( sum(item <= amount for item in amounts) / len(amounts) * 100 if amounts else None ) row_by_code = {str(row.get("ts_code") or row.get("code") or ""): row for row in index_context.get("indices") or []} value_map.update({ "index_sh_change": (row_by_code.get("000001.SH") or {}).get("pct_chg"), "index_sz_change": (row_by_code.get("399001.SZ") or {}).get("pct_chg"), "index_cy_change": (row_by_code.get("399006.SZ") or {}).get("pct_chg"), }) def missing_value(key: str) -> bool: value = value_map.get(key) return value is None or (isinstance(value, str) and not value.strip()) invalid_fields = { line_number: {key for key in keys if missing_value(key)} for line_number, keys in required.items() } if stock_common: invalid_fields[1].update(required[1]) invalid_fields[2].update(required[2]) else: if intraday and stock.get("turnover_source") in {None, "", "unavailable"}: invalid_fields[1].add("stock_turnover_relative") if intraday and stock.get("activity_source") in {None, "", "unavailable"}: invalid_fields[1].add("stock_volume_activity_ratio") if sector_common: invalid_fields[3].update(required[3]) invalid_fields[4].update(required[4]) else: if not sector.get("inner_precise", sector.get("precise")) or sector_coverage < 90: invalid_fields[3].update(key for key in required[3] if key != "sector_name") if sector.get("realtime") and not sector.get("relative_turnover"): invalid_fields[3].add("sector_relative_turnover") # The official SW index supplies only the sector's external change. A valid # membership name and member-stock leader remain usable when that quote fails. if not sector.get("outer_precise", sector.get("precise")): invalid_fields[4].add("sector_change") if index_issues: invalid_fields[6].update(required[6]) checks = [] for line_number in range(1, 7): manual_keys = [key for key in required[line_number] if key in manual_data] unresolved_fields = [ key for key in required[line_number] if key in invalid_fields[line_number] and key not in manual_data ] hard_missing_identity = line_number in {1, 2} and not stock.get("code") passed = not hard_missing_identity and not unresolved_fields status = "manual" if passed and manual_keys else "passed" if passed else "failed" reasons = [] if passed else [ *( ["请先输入并载入股票代码或名称"] if hard_missing_identity else automatic_issues[line_number] ), *( ["需补充:" + "、".join(schema[key]["label"] for key in unresolved_fields)] if unresolved_fields else [] ), ] score = float(scores[line_number - 1]["score"]) position, layer, formula = names[line_number] checks.append({ "line": line_number, "position": position, "layer": layer, "formula": formula, "status": status, "passed": passed, "reasons": reasons, "score": round(score, 3) if passed else None, "line_value": _score_to_line(score) if passed else None, "evidence": scores[line_number - 1]["evidence"] if passed else [], "fields": [ { "key": key, "label": schema[key]["label"], "unit": schema[key].get("unit", ""), "type": schema[key].get("type", "number"), "options": schema[key].get("options", []), "value": value_map.get(key), "manual": key in manual_data, "required": True, "min": schema[key].get("min"), "max": schema[key].get("max"), "integer": bool(schema[key].get("integer")), } for key in required[line_number] ], }) return checks def _resolve_heaven_stock_code(self, query: str) -> str: raw = validate_text(query, "股票代码或名称", 30, required=True) code_match = re.fullmatch(r"(\d{6})(?:\.(?:SH|SZ|BJ))?", raw.upper()) if code_match: return validate_stock_code(code_match.group(1)) candidates = self.database.search_stock_master(raw) exact = [item for item in candidates if str(item.get("name") or "").casefold() == raw.casefold()] if not exact and self.configured: try: rows = TushareClient(self.token).query( "stock_basic", {"name": raw, "list_status": "L"}, "ts_code,symbol,name,industry,market,list_date", ) except TushareError: rows = [] if rows: self.database.upsert_stock_master(rows) candidates = self.database.search_stock_master(raw) exact = [ item for item in candidates if str(item.get("name") or "").casefold() == raw.casefold() ] matches = exact or candidates if len(matches) == 1: return validate_stock_code(str(matches[0].get("code") or "")) if len(matches) > 1: choices = "、".join( f"{item.get('name') or '--'}({item.get('code') or '--'})" for item in matches[:5] ) raise ValueError(f"匹配到多只股票:{choices}。请输入六位股票代码。") raise ValueError(f"未找到股票“{raw}”,请检查名称或输入六位股票代码。") def heaven_setup( self, trade_date: str, sector_name: str = "", stock_code: str = "", manual_data: dict[str, Any] | None = None, ) -> dict[str, Any]: normalized_date = normalize_date(trade_date) dashboard = self.get_dashboard(normalized_date) data_date = normalize_date(str(dashboard.get("meta", {}).get("trade_date") or normalized_date)) recent_history = self.database.snapshot_summaries(data_date, 10) market_mode = self._heaven_market_mode(data_date, dashboard) manual_data = self._validate_heaven_manual_data(manual_data, market_mode) index_context = self._heaven_index_context(data_date, dashboard, market_mode) external_stock = None normalized_stock_code = "" if stock_code.strip(): normalized_stock_code = self._resolve_heaven_stock_code(stock_code) external_stock = self._heaven_stock_context( normalized_stock_code, data_date, dashboard, market_mode, ) external_sector = None if normalized_stock_code and self.configured: external_sector = self._heaven_sector_context( normalized_stock_code, data_date, market_mode, ) if external_sector and external_stock: external_stock["sector"] = external_sector.get("name") or external_stock.get("sector") dashboard, index_context, external_sector, external_stock = self._apply_heaven_manual_data( dashboard, index_context, external_sector, external_stock, manual_data, market_mode, data_date, normalized_stock_code, ) if external_sector and external_stock: external_stock["sector"] = external_sector.get("name") or external_stock.get("sector") sector_input = str((external_sector or {}).get("name") or sector_name.strip()) if not normalized_stock_code: data_checks = [] chart = { "available": False, "selection_required": True, "data_trade_date": data_date, "sector": "", "sector_code": "", "sector_taxonomy": "", "stock": {"code": "", "name": "", "status": ""}, "quality": { "status": "awaiting_selection", "issues": [], "principle": "", "sources": [], }, "index_context": index_context, } else: data_checks = self._heaven_line_checks( data_date, dashboard, recent_history, index_context, external_sector or {}, external_stock or {}, market_mode, manual_data, ) quality_issues = [ f"{check['position']}·{check['layer']}:{';'.join(check['reasons'])}" for check in data_checks if not check["passed"] ] if quality_issues: chart = { "available": False, "selection_required": False, "data_trade_date": data_date, "sector": str((external_sector or {}).get("name") or sector_input or "--"), "sector_code": str((external_sector or {}).get("code") or ""), "sector_taxonomy": str((external_sector or {}).get("taxonomy") or ""), "stock": { "code": normalized_stock_code, "name": str((external_stock or {}).get("name") or "--"), "status": str((external_stock or {}).get("status") or ""), }, "quality": { "status": "blocked", "issues": quality_issues, "principle": "六爻任一层缺少同日、同口径的有效数据,本系统不成卦。", "sources": self._heaven_trend_sources( data_date, index_context, external_sector, external_stock ), }, "index_context": index_context, } else: chart = build_market_hexagram( dashboard, recent_history, index_context, sector_input, normalized_stock_code, external_stock, external_sector, ) chart["available"] = True chart["selection_required"] = False manual_active = any(check["status"] == "manual" for check in data_checks) chart["quality"] = { "status": "manual" if manual_active else "verified", "issues": [], "principle": ( "自动行情与用户补充数据均已通过同一套量化公式校验。" if manual_active else "指数、板块、个股均已通过同日同口径校验。" ), "sources": [ *self._heaven_trend_sources( data_date, index_context, external_sector, external_stock ), *([{ "lines": "补录爻位", "layer": "用户补充", "realtime": market_mode == "intraday", "detail": str(manual_data.get("note") or "量化数据经原公式重新计算"), }] if manual_active else []), ], } chart["data_checks"] = data_checks chart["manual_data"] = manual_data sector_phase_overrides = self.database.list_sector_phase_overrides() field = build_five_phase_field( normalized_date, sector_phase_overrides, ) personal_profile = self.account_personal_field( normalized_date, field, public=True, ) return { "trade_date": data_date, "calendar_date": normalized_date, "market_mode": market_mode, "chart": chart, "field": field, "personal_profile": personal_profile, "sector_phase_overrides": [ {"name": name, "element": element} for name, element in sector_phase_overrides.items() ], "llm": { "configured": self.llm_configured, "model": self.llm_primary_model if self.llm_configured else "", "fallback_configured": self.llm_fallback_configured, "fallback_model": self.llm_fallback_model if self.llm_fallback_configured else "", }, } def _heaven_stock_context( self, stock_code: str, trade_date: str, dashboard: dict[str, Any], market_mode: str, ) -> dict[str, Any]: """Return the only stock contract accepted by heaven trend.""" pool_row = next( ( dict(row) for key in ("limits", "broken", "down_limits") for row in dashboard.get(key) or [] if str(row.get("code") or "") == stock_code ), {}, ) if market_mode == "intraday": if self.configured: try: quote = TushareClient(self.token).realtime_stock_quote( tushare_code(stock_code), trade_date, ) return { **quote, "status": pool_row.get("status") or "普通", "seal_amount_million": pool_row.get("seal_amount_million") or 0, "open_times": pool_row.get("open_times") or 0, "streak": pool_row.get("streak") or 0, "precise": True, } except TushareError: pass if pool_row: return { **pool_row, "data_source": "dashboard_rt" if dashboard.get("meta", {}).get("realtime") else "dashboard", "trade_date": trade_date, "realtime": bool(dashboard.get("meta", {}).get("realtime")), "precise": False, } return { "code": stock_code, "name": "--", "sector": "其他", "trade_date": trade_date, "realtime": False, "precise": False, } detail = self.get_stock_detail(stock_code, trade_date, force=True) detail_meta = detail.get("meta") or {} stock = detail.get("stock") or {} resolved_date = normalize_date(str(detail_meta.get("trade_date") or trade_date)) source = str(detail_meta.get("source") or "") return { "code": stock_code, "name": stock.get("name") or pool_row.get("name") or "--", "sector": stock.get("industry") or pool_row.get("sector") or "其他", "status": pool_row.get("status") or "普通", "change": stock.get("change") or 0, "turnover_rate": stock.get("turnover_rate") or 0, "amount_billion": stock.get("amount_billion") or 0, "seal_amount_million": pool_row.get("seal_amount_million") or 0, "open_times": pool_row.get("open_times") or 0, "streak": pool_row.get("streak") or 0, "data_source": source, "trade_date": resolved_date, "realtime": False, "precise": source == "tushare" and resolved_date == trade_date, } @staticmethod def _heaven_market_mode( trade_date: str, dashboard: dict[str, Any], now: datetime | None = None, ) -> str: """区分盘中、今日收盘和历史,避免把 rt_k 数据来源误当成交易状态。""" now = now or datetime.now().astimezone() if trade_date != now.strftime("%Y%m%d"): return "historical" meta = dashboard.get("meta") or {} status = str(meta.get("market_status") or "").lower() local_time = now.time().replace(tzinfo=None) if status == "closed" or local_time > datetime.strptime("15:05", "%H:%M").time(): return "closed" if status in {"trading", "auction", "pre_open"} or ( bool(meta.get("realtime")) and local_time >= datetime.strptime("09:15", "%H:%M").time() ): return "intraday" return "historical" @staticmethod def _heaven_trend_sources( trade_date: str, index_context: dict[str, Any], sector: dict[str, Any] | None, stock: dict[str, Any] | None, ) -> list[dict[str, Any]]: sector = sector or {} stock = stock or {} return [ { "lines": "五爻、上爻", "layer": "指数", "source": index_context.get("source") or "unavailable", "trade_date": index_context.get("trade_date") or "", "realtime": bool(index_context.get("realtime")), "detail": f"三大指数 {len(index_context.get('indices') or [])}/3", }, { "lines": "三爻、四爻", "layer": "行业", "source": sector.get("source") or "unavailable", "trade_date": sector.get("trade_date") or "", "realtime": bool(sector.get("realtime")), "detail": ( f"申万二级 {sector.get('name') or '--'} {sector.get('code') or '--'} " f"成分覆盖 {int(sector.get('quote_count') or 0)}/{int(sector.get('member_count') or 0)}" ), }, { "lines": "初爻、二爻", "layer": "个股", "source": stock.get("data_source") or "unavailable", "trade_date": stock.get("trade_date") or trade_date, "realtime": bool(stock.get("realtime")), "detail": ( f"{stock.get('name') or '--'};换手基准 " f"{stock.get('capital_trade_date') or '--'}" ), }, ] @staticmethod def _heaven_trend_quality_issues( trade_date: str, dashboard: dict[str, Any], index_context: dict[str, Any], sector: dict[str, Any] | None, stock: dict[str, Any] | None, market_mode: str = "historical", ) -> list[str]: issues: list[str] = [] intraday = market_mode == "intraday" closed = market_mode == "closed" if intraday: meta = dashboard.get("meta") or {} market_status = str(meta.get("market_status") or "") now = datetime.now().astimezone() try: updated_at = datetime.fromisoformat(str(meta.get("updated_at") or "")) if updated_at.tzinfo is None: updated_at = updated_at.replace(tzinfo=now.tzinfo) snapshot_age = (now - updated_at.astimezone(now.tzinfo)).total_seconds() except ValueError: snapshot_age = float("inf") if market_status in {"trading", "auction", "pre_open"} and snapshot_age > 120: issues.append("主行情快照超过2分钟,请点击顶部刷新") # 收盘后不再用 dashboard.market_status 作为阻断条件。盘后同步可能将 # rt_k 快照替换成同日盘后日线而不带该字段;六爻数据本身的日期、 # 完整性和来源校验已足以判断是否可以成卦。 index_date = str(index_context.get("trade_date") or "").replace("-", "") index_rows = list(index_context.get("indices") or []) index_row_dates = { str(row.get("trade_date") or "").replace("-", "") for row in index_rows } if not index_context.get("precise") or len(index_rows) < 3: issues.append("指数层缺少三大指数的有效行情") elif index_date != trade_date or index_row_dates != {trade_date}: issues.append("指数行情与目标交易日不一致") elif intraday and not index_context.get("realtime"): issues.append("盘中指数层缺少可核验的实时行情") elif not intraday and ( index_context.get("realtime") or str(index_context.get("source") or "") != "tushare" ): issues.append("历史/收盘指数层必须使用 Tushare 官方指数日线") sector = sector or {} sector_date = str(sector.get("trade_date") or "").replace("-", "") sector_coverage = float(sector.get("coverage") or 0) if not sector: issues.append("行业层缺少申万二级行业归属") elif sector.get("taxonomy") != "sw_l2": issues.append("行业层必须使用申万二级行业分类") elif sector_date != trade_date: issues.append("行业行情与目标交易日不一致") elif intraday and not sector.get("realtime"): issues.append("盘中行业层缺少申万实时行情") elif market_mode == "historical" and sector.get("realtime"): issues.append("历史行业层不能使用实时快照") elif closed and sector.get("realtime") and not sector.get("finalized"): issues.append("收盘行业层缺少15:00最终快照") if not sector.get("inner_precise", sector.get("precise")): issues.append("行业内核缺少可核验的成分行情") if not sector.get("outer_precise", sector.get("precise")): issues.append("行业外显缺少申万官方行情") if sector and sector_coverage < 90: issues.append("行业成分行情覆盖率不足90%") if sector.get("realtime") and not sector.get("relative_turnover"): issues.append("行业内核缺少相对全市场换手活跃度") stock = stock or {} stock_date = str(stock.get("trade_date") or "").replace("-", "") if not stock or not stock.get("code"): issues.append("个股层尚未载入有效标的") elif not stock.get("precise"): issues.append("个股层缺少可核验的行情数据") elif stock_date != trade_date: issues.append("个股行情与目标交易日不一致") elif intraday and not stock.get("realtime"): issues.append("盘中个股层不是 rt_k 实时行情") elif not intraday and ( stock.get("realtime") or str(stock.get("data_source") or "") != "tushare" ): issues.append("历史/收盘个股层必须使用 Tushare 官方日线") if intraday and stock and not stock.get("turnover_source"): issues.append("个股内核缺少可核验的实时换手率") elif intraday and stock.get("turnover_source") == "unavailable": issues.append("个股内核缺少流通股本,无法计算实时换手率") if intraday and stock.get("activity_source") == "unavailable": issues.append("个股内核缺少近5日量能基准") elif intraday and not stock.get("activity_source"): issues.append("个股内核缺少同时间进度量能") return issues def heaven_personal(self, payload: dict[str, Any]) -> dict[str, Any]: trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) field = build_five_phase_field( trade_date, self.database.list_sector_phase_overrides(), ) personal = self.account_personal_field(trade_date, field, public=True) if not personal: raise ValueError("请先在账号设置中保存个人命理资料。") return personal def heaven_hexagram(self, raw_lines: Any) -> dict[str, Any]: if not isinstance(raw_lines, list): raise ValueError("六爻起卦结果格式不正确。") try: lines = [int(value) for value in raw_lines] except (TypeError, ValueError) as exc: raise ValueError("六爻必须由六、七、八、九组成。") from exc return hexagram_from_lines(lines) def heaven_interpret(self, payload: dict[str, Any]) -> dict[str, Any]: mode = str(payload.get("mode") or "").strip() if mode not in {"trend", "fortune", "heart"}: raise ValueError("问天解读模式不正确。") trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) if mode in {"trend", "fortune"}: setup = self.heaven_setup( trade_date, str(payload.get("sector") or ""), str(payload.get("stock_code") or ""), payload.get("manual_data"), ) if mode == "trend": chart = setup["chart"] if not chart.get("available"): issues = ";".join((chart.get("quality") or {}).get("issues") or []) raise ValueError(f"观势数据未通过六爻校验,暂不解势:{issues}") hexagram_context = json.loads(json.dumps(chart["hexagram"], ensure_ascii=False)) for line in hexagram_context.get("lines", []): line.pop("evidence", None) line.pop("score", None) line.pop("talent", None) line.pop("layer", None) line.pop("role", None) if not line.get("moving"): line.pop("text", None) line.pop("image", None) line.pop("line_name", None) context = { "data_trade_date": setup["trade_date"], "selected_focus": { "sector": chart.get("sector") or "", "stock": chart.get("stock") or {}, }, "hexagram": hexagram_context, "movement": chart.get("movement") or {}, } else: personal_profile = self.account_personal_field( setup["calendar_date"], setup["field"], public=False, ) fortune_field = json.loads(json.dumps(setup["field"], ensure_ascii=False)) catalog = fortune_field.pop("sector_catalog", []) dominant_elements = { item.get("element") for item in fortune_field.get("balance", [])[:2] } fortune_field["industry_affinity"] = [ { "element": group.get("element"), "examples": [ item.get("name") for item in group.get("industries", [])[:8] if item.get("name") ], } for group in catalog if group.get("element") in dominant_elements ] context = { "calendar_date": setup["calendar_date"], "five_phase_field": fortune_field, "personal_profile": personal_profile, } else: context = { "hexagram": self.heaven_hexagram(payload.get("lines")), "ritual": "用户已完成30秒静心、六次三枚铜钱起卦,并在心中察看第一念。问题未输入。", } result, compiler = self._call_heaven_agent(mode, context) return { **result, "mode": mode, "compiler": compiler, "notice": "主模型调用失败,已自动切换辅助模型。" if compiler == "fallback" else "", } def _call_heaven_agent(self, mode: str, context: dict[str, Any]) -> tuple[dict[str, Any], str]: source = self.ensure_llm_access(f"heaven_{mode}") primary_error = "" if self.llm_configured: try: result = interpret_heaven( mode, context, self.llm_primary_api_key, self.llm_primary_base_url, self.llm_primary_model, ) self.record_llm_usage( f"heaven_{mode}", source, str(result.get("model") or ""), "success", int(result.get("latency_ms") or 0), ) return result, "primary" except HeavenAgentError as exc: primary_error = str(exc) if self.llm_fallback_configured: try: result = interpret_heaven( mode, context, self.llm_fallback_api_key, self.llm_fallback_base_url, self.llm_fallback_model, ) self.record_llm_usage( f"heaven_{mode}", source, str(result.get("model") or ""), "success", int(result.get("latency_ms") or 0), ) return result, "fallback" except HeavenAgentError as exc: self.record_llm_usage( f"heaven_{mode}", source, self.llm_fallback_model, "failed" ) raise ValueError( f"主模型和辅助模型均不可用。主模型:{primary_error or '未配置'};" f"辅助模型:{exc}" ) from exc self.record_llm_usage(f"heaven_{mode}", source, self.llm_primary_model, "failed") raise ValueError(f"主模型不可用且未配置辅助模型:{primary_error}") def _heaven_index_context( self, trade_date: str, dashboard: dict[str, Any], market_mode: str = "historical", ) -> dict[str, Any]: cached = self.database.get_data_snapshot("heaven_indices", trade_date) cached_valid = False if cached: cached_rows = list(cached.get("indices") or []) cached_dates = { str(row.get("trade_date") or "").replace("-", "") for row in cached_rows } cached_valid = ( len(cached_rows) == 3 and cached_dates == {trade_date} and bool(cached.get("precise")) and not cached.get("realtime") and str(cached.get("source") or "") == "tushare" and int(cached.get("schema_version") or 0) >= 3 ) if market_mode != "intraday" and cached_valid: return cached if not self.configured: error = "Tushare Token 未配置" else: try: client = TushareClient(self.token) if market_mode == "intraday": payload = self._aggregate_index_context(trade_date) payload["schema_version"] = 3 return payload payload = client.market_indices(trade_date) payload["schema_version"] = 3 if market_mode == "closed": payload["finalized"] = True self.database.save_data_snapshot( "heaven_indices", trade_date, str(payload.get("source") or "tushare"), payload, ) return payload except Exception as exc: error = str(exc) overview = dashboard.get("overview") or {} up_count = float(overview.get("up_count") or 0) down_count = float(overview.get("down_count") or 0) breadth = (up_count - down_count) / max(up_count + down_count, 1) return { "source": "market_breadth_proxy", "trade_date": trade_date, "realtime": False, "precise": False, "schema_version": 3, "notice": f"指数数据不可用,当前以市场宽度代理:{error}", "indices": [], "aggregate": { "average_pct_chg": round(breadth * 2.5, 3), "average_return_5d": 0, "average_return_20d": 0, }, } def _aggregate_index_context( self, trade_date: str, tushare_error: str = "", ) -> dict[str, Any]: quotes = self.realtime_aggregator.tencent_indices() epochs = [int(item.get("quote_time_epoch") or 0) for item in quotes] quote_dates = { datetime.fromtimestamp(epoch).astimezone().strftime("%Y%m%d") for epoch in epochs if epoch } if len(quotes) != 3 or quote_dates != {trade_date}: raise ValueError("腾讯三大指数日期与目标交易日不一致") now = datetime.now().astimezone() max_skew = 120 if now.hour >= 15 else 15 if max(epochs) - min(epochs) > max_skew: raise ValueError(f"腾讯三大指数时间差超过{max_skew}秒") code_map = { "000001": "000001.SH", "399001": "399001.SZ", "399006": "399006.SZ", } client = TushareClient(self.token) indices = [] start_date = ( datetime.strptime(trade_date, "%Y%m%d") - timedelta(days=20) ).strftime("%Y%m%d") for quote in quotes: ts_code = code_map[str(quote.get("code") or "")] history = client.query( "index_daily", {"ts_code": ts_code, "start_date": start_date, "end_date": trade_date}, "ts_code,trade_date,close,pct_chg", ) history.sort(key=lambda item: str(item.get("trade_date") or "")) completed_closes = [ float(item.get("close") or 0) for item in history if str(item.get("trade_date") or "") < trade_date and float(item.get("close") or 0) > 0 ] close_5d = ( completed_closes[-5] if len(completed_closes) >= 5 else completed_closes[0] if completed_closes else 0 ) close = float(quote.get("price") or 0) indices.append( { "ts_code": ts_code, "name": quote.get("name") or ts_code, "trade_date": trade_date, "close": close, "pct_chg": round(float(quote.get("change") or 0), 3), "return_5d": round((close / close_5d - 1) * 100, 3) if close_5d else 0, "return_20d": 0, "amount_billion": float(quote.get("amount_billion") or 0), "quote_time": quote.get("quote_time") or "", } ) return { "trade_date": trade_date, "source": "+".join( sorted({str(item.get("source") or "web_quote") for item in quotes}) + ["tushare_index_daily"] ), "realtime": True, "precise": True, "indices": indices, "aggregate": { "average_pct_chg": round( sum(item["pct_chg"] for item in indices) / len(indices), 3 ), "average_return_5d": round( sum(item["return_5d"] for item in indices) / len(indices), 3 ), "average_return_20d": 0, }, "quote_time_skew_seconds": max(epochs) - min(epochs), "notice": ( "指数实时行情来自腾讯行情,5日趋势来自Tushare历史指数。" + (f" Tushare实时指数未使用:{tushare_error}" if tushare_error else "") ), } def _heaven_sector_context( self, identifier: str, trade_date: str, market_mode: str = "historical", ) -> dict[str, Any] | None: """Return the Shenwan L2 sector context for heaven trend. 观势行业层只使用申万二级行业。外显盘中使用 rt_sw_k、历史使用 sw_daily;内核独立使用目标日期成分股行情聚合。收盘过渡期在 sw_daily 入库前接受同日15:00后的 rt_sw_k 收盘快照。 """ cache_key = f"{trade_date}:{identifier.strip().lower()}" cached = self.database.get_data_snapshot("heaven_sector", cache_key) cached_date = str((cached or {}).get("trade_date") or "").replace("-", "") cached_valid = bool( cached and cached_date == trade_date and cached.get("taxonomy") == "sw_l2" and cached.get("inner_precise", cached.get("precise")) and cached.get("outer_precise", cached.get("precise")) and not cached.get("realtime") and int(cached.get("schema_version") or 0) >= 4 ) if market_mode != "intraday" and cached_valid: return cached if not self.configured: return None try: payload = TushareClient(self.token).sw_sector_snapshot( tushare_code(identifier), trade_date, realtime_expected=market_mode == "intraday", allow_realtime_close=market_mode == "closed", ) except TushareError as exc: if cached_valid: return cached return { "name": "", "code": "", "taxonomy": "sw_l2", "source": "tushare", "trade_date": trade_date, "realtime": market_mode == "intraday", "precise": False, "inner_precise": False, "outer_precise": False, "coverage": 0, "member_count": 0, "quote_count": 0, "error": f"申万二级行业数据获取失败:{exc}", } if not payload.get("realtime") and payload.get("precise"): self.database.save_data_snapshot( "heaven_sector", cache_key, str(payload.get("source") or "tushare"), payload, ) return payload @staticmethod def _validate_mentor_history(raw_history: Any) -> list[dict[str, str]]: if not isinstance(raw_history, list): raise ValueError("问师对话历史格式不正确。") history = [] total_length = 0 for item in raw_history[-12:]: if not isinstance(item, dict) or item.get("role") not in {"user", "assistant"}: raise ValueError("问师对话历史包含无效消息。") content = str(item.get("content") or "").strip() if not content or len(content) > 5000: raise ValueError("问师对话历史消息为空或过长。") total_length += len(content) if total_length > 24_000: raise ValueError("问师对话历史过长,请清空后重新提问。") history.append({"role": item["role"], "content": content}) return history def _build_mentor_context(self, trade_date: str, question: str) -> dict[str, Any]: dashboard = self.get_dashboard(trade_date) data_trade_date = normalize_date( str(dashboard.get("meta", {}).get("trade_date") or trade_date) ) regime = self.screener.detect_regime(data_trade_date) limits = list(dashboard.get("limits") or []) broken = list(dashboard.get("broken") or []) down_limits = list(dashboard.get("down_limits") or []) yesterday_limits = list(dashboard.get("yesterday_limits") or []) all_stocks = limits + broken + down_limits + yesterday_limits matched_rows = [] codes = re.findall(r"(?= 2 and name in question): if not any(item.get("code") == code for item in matched_rows): matched_rows.append(row) stock_details = [] for code in codes[:2]: try: detail = self.get_stock_detail(code, data_trade_date) stock_details.append( { "stock": detail.get("stock") or {}, "moneyflow": detail.get("moneyflow") or {}, "recent_prices": (detail.get("prices") or [])[-20:], } ) except Exception as exc: stock_details.append({"code": code, "error": str(exc)}) dragon_tiger = None if codes or any(keyword in question for keyword in ("龙虎榜", "席位", "机构", "游资")): try: dragon_payload = self.get_dragon_tiger(data_trade_date) rows = list(dragon_payload.get("rows") or []) matched_dragon = [row for row in rows if str(row.get("code") or "") in codes] leading_dragon = sorted( rows, key=lambda row: abs(float(row.get("net_buy_million") or 0)), reverse=True, )[:12] dragon_tiger = { "summary": dragon_payload.get("summary") or {}, "matched": matched_dragon, "largest_net_flows": leading_dragon, } except Exception as exc: dragon_tiger = {"error": str(exc)} return { "data_trade_date": data_trade_date, "source": dashboard.get("meta", {}).get("source"), "notice": dashboard.get("meta", {}).get("notice") or "", "overview": dashboard.get("overview") or {}, "market_regime": regime, "recent_market_history": self.database.snapshot_summaries(data_trade_date, 10), "limit_ladder": dashboard.get("ladders") or [], "limit_performance": dashboard.get("limit_performance") or [], "hot_sectors": (dashboard.get("sectors") or [])[:20], "sector_rotation": (dashboard.get("sector_rotation") or [])[:20], "limit_up_stocks": sorted( limits, key=lambda row: ( float(row.get("streak") or 0), float(row.get("amount_billion") or 0), ), reverse=True, )[:30], "broken_stocks": sorted( broken, key=lambda row: float(row.get("amount_billion") or 0), reverse=True, )[:20], "limit_down_stocks": sorted( down_limits, key=lambda row: float(row.get("amount_billion") or 0), reverse=True, )[:25], "yesterday_limit_performance": sorted( yesterday_limits, key=lambda row: float(row.get("change") or 0), reverse=True, )[:25], "question_matched_stocks": matched_rows[:10], "stock_details": stock_details, "dragon_tiger": dragon_tiger, } def run_screener(self, payload: dict[str, Any]) -> dict[str, Any]: trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat())) regime = str(payload.get("regime") or "") if regime not in REGIMES: raise ValueError("市场阶段不支持。") strategy_name = validate_text(payload.get("strategy_name"), "策略名称", 60, required=True) formula = payload.get("formula") or {} realtime_snapshot = None dashboard = self.get_dashboard(trade_date) if self.configured and dashboard.get("meta", {}).get("realtime"): try: realtime_snapshot = TushareClient(self.token).realtime_factor_snapshot(trade_date) except TushareError as exc: raise ValueError(f"实时选股行情不可用,已停止筛选:{exc}") from exc return self.screener.screen( self.current_user_id, trade_date, formula, regime, strategy_name, bool(payload.get("run_backtest", True)), realtime_snapshot, ) def get_dragon_tiger(self, trade_date: str, force: bool = False) -> dict[str, Any]: normalized_date = normalize_date(trade_date) cache_kind = "hot_money_detail_v3" if not force: cached = self.database.get_data_snapshot(cache_kind, normalized_date) if ( cached and cached.get("meta", {}).get("source") == "tushare" and cached.get("meta", {}).get("status") == "success" and int(cached.get("meta", {}).get("schema_version") or 0) == 3 ): cached["meta"] = {**cached.get("meta", {}), "cached": True} return cached if self.configured: try: payload = TushareClient(self.token).dragon_tiger(normalized_date) except TushareError as exc: return { "meta": { "requested_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", "trade_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", "source": "tushare_error", "status": "error", "schema_version": 3, "cached": False, "updated_at": datetime.now().astimezone().isoformat(timespec="seconds"), "notice": f"游资接口不可用:{exc}", }, "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": [], } payload["meta"]["cached"] = False if payload.get("meta", {}).get("status") == "success": self.database.save_data_snapshot(cache_kind, normalized_date, "tushare", payload) return payload return { "meta": { "requested_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", "trade_date": f"{normalized_date[:4]}-{normalized_date[4:6]}-{normalized_date[6:8]}", "source": "unavailable", "status": "unavailable", "schema_version": 3, "cached": False, "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": [], } def _search_market_directory(self) -> list[dict[str, Any]]: cached = self.database.get_data_snapshot("search_directory", "ths") or {} cached_items = list(cached.get("items") or []) if cached_items and int(cached.get("schema_version") or 0) >= 2: return cached_items if not self.configured: return cached_items try: rows = TushareClient(self.token).query( "ths_index", {}, "ts_code,name,count,exchange,list_date,type", ) except TushareError: return cached_items items = [] for row in rows: mapping = THS_SEARCH_TYPES.get(str(row.get("type") or "").upper()) code = str(row.get("ts_code") or "").strip().upper() name = str(row.get("name") or "").strip() if not mapping or not code or not name or str(row.get("exchange") or "").upper() != "A": continue entity_type, subtitle = mapping items.append( { "id": code, "code": code, "name": name, "type": entity_type, "subtitle": subtitle, "member_count": int(float(row.get("count") or 0)), } ) if items: self.database.save_data_snapshot( "search_directory", "ths", "tushare", {"schema_version": 2, "items": items} ) return items @staticmethod def _search_match_score(item: dict[str, Any], query: str) -> tuple[int, int, str]: name = str(item.get("name") or "").casefold() code = str(item.get("code") or item.get("id") or "").casefold() needle = query.casefold() if code == needle: rank = 0 elif name == needle: rank = 1 elif code.startswith(needle): rank = 2 elif name.startswith(needle): rank = 3 else: rank = 4 return rank, len(name), code def search_entities(self, query: str, trade_date: str) -> dict[str, Any]: needle = str(query or "").strip() normalized_date = normalize_date(trade_date) groups: dict[str, list[dict[str, Any]]] = { "stocks": [], "sectors": [], "themes": [], "indices": [], } if not needle: return {"query": "", "trade_date": normalized_date, "groups": groups} stocks = [] for row in self.database.search_stock_master(needle, 12): stocks.append( { "id": str(row.get("code") or ""), "code": str(row.get("code") or ""), "name": str(row.get("name") or "--"), "type": "stock", "type_label": SEARCH_TYPE_LABELS["stock"], "industry": str(row.get("industry") or "其他"), "market": str(row.get("market") or ""), "subtitle": " · ".join( part for part in (str(row.get("industry") or ""), str(row.get("market") or "")) if part ) or "A股", } ) groups["stocks"] = stocks[:8] market_items = list(self._search_market_directory()) + [dict(item) for item in SEARCH_INDEXES] matched = [ item for item in market_items if needle.casefold() in str(item.get("name") or "").casefold() or needle.casefold() in str(item.get("code") or "").casefold() ] matched.sort(key=lambda item: self._search_match_score(item, needle)) group_keys = {"sector": "sectors", "theme": "themes", "index": "indices"} for item in matched: group_key = group_keys.get(str(item.get("type") or "")) if not group_key or len(groups[group_key]) >= 8: continue groups[group_key].append( { **item, "type_label": SEARCH_TYPE_LABELS[str(item["type"])], } ) return {"query": needle, "trade_date": normalized_date, "groups": groups} def get_search_detail( self, entity_type: str, identifier: str, trade_date: str ) -> dict[str, Any]: entity_type = str(entity_type or "").strip().lower() identifier = str(identifier or "").strip().upper() normalized_date = normalize_date(trade_date) if entity_type not in {"sector", "theme", "index"}: raise ValueError("搜索详情类型不支持。") if not re.fullmatch(r"[A-Z0-9.]{3,24}", identifier): raise ValueError("搜索详情标识无效。") if not self.configured: raise ValueError("行情数据源尚未配置。") if entity_type == "index": index_basic = next((item for item in SEARCH_INDEXES if item["id"] == identifier), None) if not index_basic: raise ValueError("暂不支持该指数详情。") return self._index_search_detail(index_basic, normalized_date) directory = self._search_market_directory() basic = next( ( item for item in directory if item.get("id") == identifier and item.get("type") == entity_type ), None, ) if not basic: raise ValueError("未找到对应的板块或题材。") return self._ths_search_detail(basic, normalized_date) def _ths_search_detail( self, basic: dict[str, Any], trade_date: str ) -> dict[str, Any]: client = TushareClient(self.token) resolved_date, _ = client.resolve_trade_context(trade_date) end = datetime.strptime(resolved_date, "%Y%m%d") start_date = (end - timedelta(days=190)).strftime("%Y%m%d") identifier = str(basic["id"]) snapshot = client.sector_snapshot(identifier, resolved_date) rows = client.query( "ths_daily", {"ts_code": identifier, "start_date": start_date, "end_date": resolved_date}, "ts_code,trade_date,open,high,low,close,pct_change,vol,turnover_rate,total_mv,float_mv", ) rows.sort(key=lambda item: str(item.get("trade_date") or "")) series = [ { "trade_date": self._display_compact_date(str(row.get("trade_date") or "")), "open": float(row.get("open") or 0), "high": float(row.get("high") or 0), "low": float(row.get("low") or 0), "close": float(row.get("close") or 0), "change": float(row.get("pct_change") or 0), "volume": float(row.get("vol") or 0), "turnover_rate": float(row.get("turnover_rate") or 0), } for row in rows[-90:] ] latest = series[-1] if series else {} snapshot_is_current = str(snapshot.get("trade_date") or "").replace("-", "") == resolved_date change = float( snapshot.get("change") if snapshot_is_current and snapshot.get("change") is not None else latest.get("change") or 0 ) turnover_rate = float( snapshot.get("turnover_rate") if snapshot_is_current and snapshot.get("turnover_rate") is not None else latest.get("turnover_rate") or 0 ) metrics = [ {"label": "涨跌幅", "value": round(change, 2), "unit": "%", "tone": "change"}, {"label": "换手率", "value": round(turnover_rate, 2), "unit": "%"}, {"label": "成份数量", "value": int(float(basic.get("member_count") or 0)), "unit": "只"}, ] up_count = int(float(snapshot.get("up_count") or 0)) down_count = int(float(snapshot.get("down_count") or 0)) if up_count or down_count: metrics.extend( [ {"label": "上涨家数", "value": up_count, "unit": "家"}, {"label": "下跌家数", "value": down_count, "unit": "家"}, ] ) leader = str(snapshot.get("leader") or "").strip() if leader and leader != "--": metrics.extend( [ {"label": "领涨标的", "value": leader, "unit": ""}, {"label": "领涨幅", "value": round(float(snapshot.get("leading_pct") or 0), 2), "unit": "%", "tone": "change"}, ] ) return { "meta": { "trade_date": self._display_compact_date(resolved_date), "realtime": bool(snapshot.get("realtime")), }, "entity": { "id": identifier, "code": identifier, "name": str(snapshot.get("name") or basic.get("name") or "--"), "type": str(basic.get("type") or "sector"), "type_label": SEARCH_TYPE_LABELS[str(basic.get("type") or "sector")], "subtitle": str(basic.get("subtitle") or ""), "value": float(latest.get("close") or 0), "change": change, }, "series": series, "metrics": metrics, } def _index_search_detail( self, basic: dict[str, Any], trade_date: str ) -> dict[str, Any]: client = TushareClient(self.token) resolved_date, _ = client.resolve_trade_context(trade_date) payload = ( client.realtime_market_indices(resolved_date) if client.should_use_realtime(trade_date, resolved_date) else client.market_indices(resolved_date, 90) ) current = next( (item for item in payload.get("indices") or [] if item.get("ts_code") == basic["id"]), None, ) if not current: raise ValueError("该指数暂无可用行情。") end = datetime.strptime(resolved_date, "%Y%m%d") rows = client.query( "index_daily", { "ts_code": basic["id"], "start_date": (end - timedelta(days=190)).strftime("%Y%m%d"), "end_date": resolved_date, }, "ts_code,trade_date,open,high,low,close,pct_chg,vol,amount", ) rows.sort(key=lambda item: str(item.get("trade_date") or "")) series = [ { "trade_date": self._display_compact_date(str(row.get("trade_date") or "")), "open": float(row.get("open") or 0), "high": float(row.get("high") or 0), "low": float(row.get("low") or 0), "close": float(row.get("close") or 0), "change": float(row.get("pct_chg") or 0), "volume": float(row.get("vol") or 0), } for row in rows[-90:] ] return { "meta": { "trade_date": self._display_compact_date(str(current.get("trade_date") or resolved_date)), "realtime": bool(payload.get("realtime")), }, "entity": { **basic, "type_label": SEARCH_TYPE_LABELS["index"], "value": float(current.get("close") or 0), "change": float(current.get("pct_chg") or 0), }, "series": series, "metrics": [ {"label": "涨跌幅", "value": round(float(current.get("pct_chg") or 0), 2), "unit": "%", "tone": "change"}, {"label": "近5日", "value": round(float(current.get("return_5d") or 0), 2), "unit": "%", "tone": "change"}, {"label": "近20日", "value": round(float(current.get("return_20d") or 0), 2), "unit": "%", "tone": "change"}, {"label": "成交额", "value": round(float(current.get("amount_billion") or 0), 2), "unit": "亿"}, ], } def get_stock_detail( self, code: str, trade_date: str, force: bool = False ) -> dict[str, Any]: code = validate_stock_code(code) normalized_date = normalize_date(trade_date) cache_key = f"{code}:{normalized_date}" if not force: cached = self.database.get_data_snapshot("stock_detail", cache_key) 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: try: payload = TushareClient(self.token).stock_detail( tushare_code(code), normalized_date ) if not payload.get("prices"): raise TushareError("No price history returned") except TushareError as 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 = 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) return self._enrich_stock_detail(payload) def get_stock_preview( self, code: str, trade_date: str, force: bool = False ) -> dict[str, Any]: code = validate_stock_code(code) detail = self.get_stock_detail(code, trade_date, force) detail_meta = detail.get("meta") or {} resolved_date = str(detail_meta.get("trade_date") or trade_date) compact_date = normalize_date(resolved_date) intraday_points: list[dict[str, Any]] = [] intraday_status = "unavailable" intraday_notice = "未配置 Tushare Token,分时数据不可用。" if self.configured: cache_key = f"{code}:{compact_date}" cached = None if force else self.database.get_data_snapshot("stock_intraday", cache_key) if cached and cached.get("points"): intraday_points = list(cached["points"]) intraday_status = "available" intraday_notice = "" else: try: intraday = TushareClient(self.token).stock_intraday( tushare_code(code), compact_date ) intraday_points = list(intraday.get("points") or []) if intraday_points: intraday_status = "available" intraday_notice = "" self.database.save_data_snapshot( "stock_intraday", cache_key, "tushare", intraday ) else: intraday_status = "empty" intraday_notice = "该交易日暂无分时数据。" except TushareError as exc: intraday_status = "unavailable" intraday_notice = f"Tushare 分时接口不可用:{exc}" prices = list(detail.get("prices") or [])[-60:] stock = dict(detail.get("stock") or {"code": code}) realtime = False if self.configured and compact_date == date.today().strftime("%Y%m%d"): try: quote = TushareClient(self.token).realtime_stock_quote( tushare_code(code), compact_date, ) realtime_bar = { "trade_date": f"{compact_date[:4]}-{compact_date[4:6]}-{compact_date[6:]}", "open": quote["open"], "high": quote["high"], "low": quote["low"], "close": quote["price"], "change": quote["change"], "volume": quote["volume"] / 100, "amount_billion": quote["amount_billion"], "realtime": True, } if prices and str(prices[-1].get("trade_date") or "").replace("-", "") == compact_date: prices[-1] = realtime_bar else: prices.append(realtime_bar) prices = prices[-60:] stock.update( { "name": quote["name"], "industry": quote["sector"], "price": quote["price"], "change": quote["change"], "amount_billion": quote["amount_billion"], "turnover_rate": quote["turnover_rate"], } ) realtime = True except TushareError: realtime = False return { "meta": { "trade_date": resolved_date, "source": detail_meta.get("source") or "unavailable", "notice": detail_meta.get("notice") or "", "intraday_status": intraday_status, "intraday_notice": intraday_notice, "realtime": realtime, "refresh_interval_seconds": 10 if realtime else 0, }, "stock": stock, "prices": prices, "intraday": intraday_points, } def save_reason(self, trade_date: str, code: str, reason: str) -> None: normalized_date = normalize_date(trade_date) code = validate_stock_code(code) reason = reason.strip() if not reason or len(reason) > 200: raise ValueError("涨停原因应为 1 至 200 个字符。") self.database.save_reason_override(normalized_date, code, reason) def backfill(self, start_date: str, end_date: str) -> list[dict[str, Any]]: start = datetime.strptime(normalize_date(start_date), "%Y%m%d").date() end = datetime.strptime(normalize_date(end_date), "%Y%m%d").date() if start > end: raise ValueError("开始日期不能晚于结束日期。") weekdays = [] current = start while current <= end: if current.weekday() < 5: weekdays.append(current) current += timedelta(days=1) if len(weekdays) > 15: raise ValueError("单次最多回补 15 个工作日。") results = [] for day in weekdays: dashboard = self.sync_dashboard(day.strftime("%Y%m%d")) results.append( { "requested_date": day.isoformat(), "trade_date": dashboard["meta"]["trade_date"], "source": dashboard["meta"]["source"], "records": self._record_count(dashboard), } ) return results def _stock_identity(self, code: str, trade_date: str) -> tuple[str, str]: snapshot = self.database.get_snapshot(trade_date) or {} for key in ("limits", "broken", "down_limits"): for row in snapshot.get(key) or []: if str(row.get("code")) == code: return row.get("name") or "--", row.get("sector") or "其他" for item in self.database.list_watchlist(self.current_user_id): if item["code"] == code: return item["name"], item["sector"] or "其他" return "--", "其他" def _enrich_stock_detail(self, payload: dict[str, Any]) -> dict[str, Any]: result = dict(payload) stock = dict(payload.get("stock") or {}) code = str(stock.get("code") or "") watched = { item["code"]: item for item in self.database.list_watchlist(self.current_user_id) } stock["watchlist"] = watched.get(code) result["stock"] = stock result["notes"] = self.database.list_notes(self.current_user_id, code=code) return result def _apply_reason_overrides(self, dashboard: dict[str, Any]) -> dict[str, Any]: trade_date = str(dashboard.get("meta", {}).get("trade_date", "")).replace("-", "") overrides = self.database.reason_overrides(trade_date) if not overrides: return dashboard for key in ("limits", "broken", "down_limits"): for row in dashboard.get(key) or []: if row.get("code") in overrides: row["reason"] = overrides[row["code"]] row["reason_source"] = "manual" return dashboard def _apply_seat_aliases(self, payload: dict[str, Any]) -> dict[str, Any]: aliases = self.database.list_seat_aliases() result = dict(payload) rows = payload.get("rows") or [] for row in rows: for institution in row.get("institutions") or []: institution["alias"] = aliases.get(institution.get("seat_name", ""), "") traders: dict[tuple[str, str], dict[str, Any]] = {} unclassified: dict[str, dict[str, Any]] = {} seen_operations: set[tuple[Any, ...]] = set() builtin_aliases = { "国泰海通证券股份有限公司南京太平南路证券营业部": "作手新一", } for row in rows: for institution in row.get("institutions") or []: seat_name = str(institution.get("seat_name") or "未知席位").strip() saved_alias = str(institution.get("alias") or "").strip() builtin_alias = builtin_aliases.get(seat_name, "") if saved_alias or builtin_alias: identity_name = saved_alias or builtin_alias identity_type = "trader" recognized = True identity_source = "manual" if saved_alias else "builtin" elif "机构专用" in seat_name: identity_name = "机构专用" identity_type = "institution" recognized = True identity_source = "system" elif "沪股通专用" in seat_name or "深股通专用" in seat_name: identity_name = "北向资金" identity_type = "channel" recognized = True identity_source = "system" else: identity_name = seat_name identity_type = "unclassified" recognized = False identity_source = "raw" buy = round(float(institution.get("buy_million") or 0), 2) sell = round(float(institution.get("sell_million") or 0), 2) net_buy = round(float(institution.get("net_buy_million") or 0), 2) operation_key = (row.get("code"), seat_name, buy, sell, net_buy) if operation_key in seen_operations: continue seen_operations.add(operation_key) group_key = (identity_type, identity_name) group = traders.setdefault( group_key, { "name": identity_name, "identity_type": identity_type, "identity_source": identity_source, "recognized": recognized, "buy_million": 0.0, "sell_million": 0.0, "net_buy_million": 0.0, "seat_names": set(), "stock_codes": set(), "operations": [], }, ) group["buy_million"] += buy group["sell_million"] += sell group["net_buy_million"] += net_buy group["seat_names"].add(seat_name) group["stock_codes"].add(str(row.get("code") or "")) group["operations"].append( { "code": row.get("code") or "", "name": row.get("name") or "--", "change": row.get("change") or 0, "direction": "买入" if net_buy > 0 else "卖出" if net_buy < 0 else "持平", "buy_million": buy, "sell_million": sell, "net_buy_million": net_buy, "reason": row.get("reason") or "--", "seat_name": seat_name, "seat_alias": identity_name if recognized else "", } ) if not recognized: pending = unclassified.setdefault( seat_name, { "seat_name": seat_name, "stock_codes": set(), "operation_count": 0, "buy_million": 0.0, "sell_million": 0.0, "net_buy_million": 0.0, }, ) pending["stock_codes"].add(str(row.get("code") or "")) pending["operation_count"] += 1 pending["buy_million"] += buy pending["sell_million"] += sell pending["net_buy_million"] += net_buy type_order = {"trader": 0, "institution": 1, "channel": 2, "unclassified": 3} aggregated = list(traders.values()) aggregated.sort( key=lambda item: ( type_order.get(item["identity_type"], 9), -abs(item["net_buy_million"]), item["name"], ) ) for index, group in enumerate(aggregated, start=1): group["id"] = f"identity-{index}" group["buy_million"] = round(group["buy_million"], 2) group["sell_million"] = round(group["sell_million"], 2) group["net_buy_million"] = round(group["net_buy_million"], 2) group["seat_count"] = len(group.pop("seat_names")) group["stock_count"] = len(group.pop("stock_codes")) group["operation_count"] = len(group["operations"]) group["operations"].sort( key=lambda item: abs(float(item.get("net_buy_million") or 0)), reverse=True ) pending_seats = list(unclassified.values()) for pending in pending_seats: pending["stock_count"] = len(pending.pop("stock_codes")) pending["buy_million"] = round(pending["buy_million"], 2) pending["sell_million"] = round(pending["sell_million"], 2) pending["net_buy_million"] = round(pending["net_buy_million"], 2) pending_seats.sort(key=lambda item: abs(item["net_buy_million"]), reverse=True) operation_count = sum(item["operation_count"] for item in aggregated) active_stocks = { operation["code"] for item in aggregated for operation in item["operations"] } seat_net_buy = round(sum(item["net_buy_million"] for item in aggregated), 2) result["rows"] = rows result["traders"] = aggregated result["unclassified_seats"] = pending_seats result["summary"] = { **(payload.get("summary") or {}), "trader_count": sum(item["identity_type"] == "trader" for item in aggregated), "identity_count": len(aggregated), "operation_count": operation_count, "active_stock_count": len(active_stocks), "seat_net_buy_million": seat_net_buy, "unclassified_count": len(pending_seats), } return result def _with_storage(self, dashboard: dict[str, Any], cached: bool) -> dict[str, Any]: result = dict(dashboard) result["meta"] = { **dashboard.get("meta", {}), "storage": "sqlite", "cached": cached, } return result @staticmethod def _record_count(dashboard: dict[str, Any]) -> int: return sum( len(dashboard.get(key) or []) for key in ("limits", "broken", "down_limits", "yesterday_limits") ) SERVICE = DashboardService() class RequestHandler(BaseHTTPRequestHandler): server_version = "XiaobaiReviewWeb/0.8" def do_GET(self) -> None: parsed = urlparse(self.path) if parsed.path == "/api/health": self.send_json( { "ok": True, "storage": "sqlite", "account_required": True, "time": datetime.now().astimezone().isoformat(timespec="seconds"), } ) return if parsed.path == "/api/auth/me": self.auth_me() return if parsed.path.startswith("/api/") and not self.require_auth(): return if parsed.path.startswith("/api/admin/") and not self.require_admin(): return 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( {"ok": True, **SERVICE.system_status(), "users": SERVICE.admin_users()} ) return if parsed.path == "/api/account/status": self.send_json({"ok": True, **SERVICE.status()}) return if parsed.path == "/api/dashboard": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json(SERVICE.get_dashboard(trade_date, False)) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) except Exception as exc: self.send_json({"error": f"数据加载失败:{exc}"}, HTTPStatus.INTERNAL_SERVER_ERROR) return if parsed.path == "/api/realtime-aggregate/health": query = parse_qs(parsed.query) try: self.send_json( { "ok": True, "aggregate": SERVICE.realtime_aggregate_health( query.get("sector", [""])[0] ), } ) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/sentiment/history": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: limit = int(query.get("limit", ["20"])[0]) self.send_json(SERVICE.sentiment_history(trade_date, limit)) except (TypeError, ValueError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/rotation/history": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json(SERVICE.rotation_history(trade_date, 9)) except (TypeError, ValueError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/dragon-tiger": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] force = query.get("force", ["0"])[0] == "1" try: self.send_json(SERVICE.get_dragon_tiger(trade_date, force)) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/search": query = parse_qs(parsed.query) search_query = query.get("q", [""])[0] trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json(SERVICE.search_entities(search_query, trade_date)) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/search/detail": query = parse_qs(parsed.query) entity_type = query.get("type", [""])[0] identifier = query.get("id", [""])[0] trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json( SERVICE.get_search_detail(entity_type, identifier, trade_date) ) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) except TushareError as exc: self.send_json({"error": f"行情加载失败:{exc}"}, HTTPStatus.BAD_REQUEST) return stock_preview_match = re.fullmatch(r"/api/stock/(\d{6})/preview", parsed.path) if stock_preview_match: query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] force = query.get("force", ["0"])[0] == "1" try: self.send_json( SERVICE.get_stock_preview(stock_preview_match.group(1), trade_date, force) ) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return stock_match = re.fullmatch(r"/api/stock/(\d{6})", parsed.path) if stock_match: query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] force = query.get("force", ["0"])[0] == "1" try: self.send_json(SERVICE.get_stock_detail(stock_match.group(1), trade_date, force)) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/watchlist": self.send_json( {"items": SERVICE.database.list_watchlist(SERVICE.current_user_id)} ) return if parsed.path == "/api/notes": query = parse_qs(parsed.query) code = query.get("code", [""])[0] trade_date = query.get("trade_date", [""])[0].replace("-", "") scope = query.get("scope", ["all"])[0] if scope not in {"all", "daily", "stock"}: self.send_json({"error": "复盘记录范围不支持。"}, HTTPStatus.BAD_REQUEST) return self.send_json( { "items": SERVICE.database.list_notes( SERVICE.current_user_id, code, trade_date, scope ) } ) return if parsed.path == "/api/seat-aliases": self.send_json({"items": SERVICE.database.list_seat_aliases()}) return if parsed.path == "/api/screener/setup": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json(SERVICE.screener_setup(trade_date)) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return if parsed.path == "/api/mentors/setup": query = parse_qs(parsed.query) trade_date = query.get("trade_date", [date.today().isoformat()])[0] try: self.send_json(SERVICE.mentor_setup(trade_date)) 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] sector_name = query.get("sector", [""])[0] stock_code = query.get("stock_code", [""])[0] manual_data = None manual_text = query.get("manual_data", [""])[0] if manual_text: try: manual_data = json.loads(manual_text) except json.JSONDecodeError: self.send_json({"error": "六爻补录数据格式不正确。"}, HTTPStatus.BAD_REQUEST) return try: self.send_json( SERVICE.heaven_setup( trade_date, sector_name, stock_code, manual_data, ) ) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return self.serve_static(parsed.path) def do_POST(self) -> None: parsed = urlparse(self.path) if parsed.path == "/api/auth/register": self.auth_register() return if parsed.path == "/api/auth/login": self.auth_login() return if not self.require_auth() or not self.require_csrf(): return if parsed.path.startswith("/api/admin/") and not self.require_admin(): return if parsed.path == "/api/backfill" and not self.require_admin(): return if parsed.path in { "/api/screener/sync", "/api/screener/compile", "/api/screener/strategies", "/api/screener/run", "/api/mentors/chat", "/api/heaven/hexagram", "/api/heaven/personal", "/api/heaven/interpret", "/api/heaven/sector-phases", } and not self.require_member(): return if parsed.path == "/api/auth/logout": self.auth_logout() return if parsed.path == "/api/account/birth-profile": self.save_birth_profile() return if parsed.path == "/api/account/password": self.change_password() return if parsed.path == "/api/admin/settings": self.save_system_settings() return if parsed.path == "/api/admin/settings/test": self.test_system_llm_settings() return if parsed.path == "/api/admin/membership": self.save_membership() return if parsed.path == "/api/admin/refresh": self.start_background_refresh() return if parsed.path == "/api/watchlist": self.save_watchlist() return if parsed.path == "/api/notes": 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": self.backfill_data() return if parsed.path == "/api/screener/sync": self.sync_screener_data() return if parsed.path == "/api/screener/compile": self.compile_screener_strategy() return if parsed.path == "/api/screener/strategies": self.save_screener_strategy() return if parsed.path == "/api/screener/run": self.run_screener() return if parsed.path == "/api/mentors/chat": self.mentor_chat() return if parsed.path == "/api/heaven/hexagram": self.heaven_hexagram() return if parsed.path == "/api/heaven/personal": self.heaven_personal() return if parsed.path == "/api/heaven/interpret": self.heaven_interpret() return self.send_json({"error": "Not found"}, HTTPStatus.NOT_FOUND) def do_DELETE(self) -> None: parsed = urlparse(self.path) if not self.require_auth() or not self.require_csrf(): return if parsed.path == "/api/account/birth-profile": 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(): return try: result = SERVICE.delete_screener_strategy(int(strategy_match.group(1))) self.send_json({"ok": True, **result}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) return watchlist_match = re.fullmatch(r"/api/watchlist/(\d{6})", parsed.path) if watchlist_match: deleted = SERVICE.database.delete_watchlist( SERVICE.current_user_id, watchlist_match.group(1) ) self.send_json({"ok": True, "deleted": deleted}) return note_match = re.fullmatch(r"/api/notes/(\d+)", parsed.path) if note_match: deleted = SERVICE.database.delete_note( SERVICE.current_user_id, int(note_match.group(1)) ) self.send_json({"ok": True, "deleted": deleted}) return sector_phase_match = re.fullmatch(r"/api/heaven/sector-phases/(.+)", parsed.path) if sector_phase_match: if not self.require_admin(): return name = unquote(sector_phase_match.group(1)).strip() deleted = SERVICE.database.delete_sector_phase_override(name) self.send_json({"ok": True, "deleted": deleted}) return self.send_json({"error": "Not found"}, HTTPStatus.NOT_FOUND) def auth_register(self) -> None: try: body = self.read_json_body() result = SERVICE.register_account( str(body.get("username") or ""), str(body.get("password") or ""), ) self.send_json( { "ok": True, "authenticated": True, "user": result["user"], "csrf_token": result["csrf_token"], }, HTTPStatus.CREATED, {"Set-Cookie": self.session_cookie(result["session_token"])}, ) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def auth_login(self) -> None: try: body = self.read_json_body() result = SERVICE.login_account( str(body.get("username") or ""), str(body.get("password") or ""), ) self.send_json( { "ok": True, "authenticated": True, "user": result["user"], "csrf_token": result["csrf_token"], }, headers={"Set-Cookie": self.session_cookie(result["session_token"])}, ) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.UNAUTHORIZED) def auth_me(self) -> None: if not self.require_auth(send_error=False): self.send_json( { "ok": True, "authenticated": False, "registration_required": SERVICE.database.count_users() == 0, } ) return self.send_json( { "ok": True, "authenticated": True, "user": { "id": int(self.auth_user["id"]), "username": str(self.auth_user["username"]), "role": str(self.auth_user.get("role") or "user"), "membership": SERVICE.membership(), }, "csrf_token": str(self.auth_user["csrf_token"]), } ) def auth_logout(self) -> None: raw_token = self.session_token() if raw_token: SERVICE.database.delete_session(token_hash(raw_token)) self.send_json( {"ok": True}, headers={"Set-Cookie": self.session_cookie("", clear=True)}, ) def save_birth_profile(self) -> None: try: body = self.read_json_body() personal = SERVICE.save_birth_profile(body) self.send_json({"ok": True, "personal": personal}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def change_password(self) -> None: try: body = self.read_json_body() current = str(body.get("current_password") or "") new = str(body.get("new_password") or "") confirmation = str(body.get("confirm_password") or "") if new != confirmation: raise ValueError("两次输入的新密码不一致。") SERVICE.change_password(current, new) self.send_json({"ok": True}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def session_token(self) -> str: cookie = SimpleCookie() try: cookie.load(self.headers.get("Cookie", "")) except Exception: return "" morsel = cookie.get(SESSION_COOKIE) return morsel.value if morsel else "" def require_auth(self, send_error: bool = True) -> bool: raw_token = self.session_token() user = SERVICE.database.session_user(token_hash(raw_token)) if raw_token else None if not user: if send_error: self.send_json({"error": "请先登录。"}, HTTPStatus.UNAUTHORIZED) return False self.auth_user = user SERVICE.bind_user(int(user["id"])) return True def require_csrf(self) -> bool: supplied = self.headers.get("X-CSRF-Token", "") expected = str(getattr(self, "auth_user", {}).get("csrf_token") or "") if not supplied or not secrets.compare_digest(supplied, expected): self.send_json({"error": "请求校验失败,请刷新页面后重试。"}, HTTPStatus.FORBIDDEN) return False return True def require_admin(self) -> bool: if str(getattr(self, "auth_user", {}).get("role") or "user") != "admin": self.send_json({"error": "需要管理员权限。"}, HTTPStatus.FORBIDDEN) return False return True def require_member(self) -> bool: if SERVICE.membership()["active"]: return True self.send_json( {"error": "该功能仅对有效会员开放,请联系管理员开通会员。", "code": "membership_required"}, HTTPStatus.FORBIDDEN, ) return False def session_cookie(self, value: str, clear: bool = False) -> str: max_age = 0 if clear else SESSION_MAX_AGE cookie = ( f"{SESSION_COOKIE}={value}; Path=/; HttpOnly; SameSite=Lax; Max-Age={max_age}" ) if self.headers.get("X-Forwarded-Proto", "").lower() == "https": cookie += "; Secure" return cookie def save_llm_settings(self) -> None: try: body = self.read_json_body() SERVICE.save_llm_settings( body.get("primary") or {}, body.get("fallback") or {}, bool(body.get("fallback_enabled")), ) self.send_json( { "ok": True, "configured": SERVICE.llm_configured, "model": SERVICE.llm_primary_model, "fallback_configured": SERVICE.llm_fallback_configured, "fallback_model": SERVICE.llm_fallback_model, } ) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_llm_mode(self) -> None: try: body = self.read_json_body() SERVICE.save_llm_mode(str(body.get("mode") or "auto")) self.send_json({"ok": True, "llm_access": SERVICE.llm_access_status()}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_system_settings(self) -> None: try: result = SERVICE.save_system_settings(self.read_json_body()) self.send_json({"ok": True, **result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def test_system_llm_settings(self) -> None: try: body = self.read_json_body() result = SERVICE.test_system_llm_profile( str(body.get("model_id") or ""), body.get("profile") or {} ) self.send_json({"ok": True, "result": result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_membership(self) -> None: try: SERVICE.update_membership(self.read_json_body()) self.send_json({"ok": True, "users": SERVICE.admin_users()}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def start_background_refresh(self) -> None: try: body = self.read_json_body(allow_empty=True) started = SERVICE.request_background_sync( str(body.get("trade_date") or date.today().isoformat()) ) self.send_json( { "ok": True, "started": started, "message": "后台刷新已开始" if started else "已有后台刷新任务正在运行", }, HTTPStatus.ACCEPTED, ) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def test_llm_settings(self) -> None: try: body = self.read_json_body() role = str(body.get("role") or "") profile = body.get("profile") or {} result = SERVICE.test_llm_profile(role, profile) self.send_json({"ok": True, "result": result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_watchlist(self) -> None: try: body = self.read_json_body() code = validate_stock_code(str(body.get("code", ""))) name = validate_text(body.get("name"), "股票名称", 30, required=True) sector = validate_text(body.get("sector"), "所属板块", 50) color = str(body.get("color") or "red") if color not in {"red", "blue", "green", "amber"}: raise ValueError("标记颜色不支持。") SERVICE.database.save_watchlist( SERVICE.current_user_id, code, name, sector, color ) self.send_json( { "ok": True, "items": SERVICE.database.list_watchlist(SERVICE.current_user_id), } ) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_note(self) -> None: try: body = self.read_json_body() code = str(body.get("code") or "").strip() if code: code = validate_stock_code(code) stock_name = validate_text(body.get("stock_name"), "股票名称", 30) trade_date = normalize_date(str(body.get("trade_date") or date.today().isoformat())) content = validate_text(body.get("content"), "复盘内容", 5000) plan = validate_text(body.get("plan"), "明日计划", 2000) if not content and not plan: raise ValueError("复盘内容和明日计划不能同时为空。") raw_id = body.get("id") note_id = int(raw_id) if raw_id else None saved_id = SERVICE.database.save_note( SERVICE.current_user_id, code, stock_name, trade_date, content, plan, note_id, ) self.send_json({"ok": True, "id": saved_id}) except (ValueError, TypeError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_reason(self) -> None: try: body = self.read_json_body() SERVICE.save_reason( str(body.get("trade_date") or ""), str(body.get("code") or ""), str(body.get("reason") or ""), ) self.send_json({"ok": True}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_seat_alias(self) -> None: try: body = self.read_json_body() seat_name = validate_text(body.get("seat_name"), "席位名称", 200, required=True) alias = validate_text(body.get("alias"), "席位别名", 50, required=True) SERVICE.database.save_seat_alias(seat_name, alias) self.send_json({"ok": True}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_sector_phase_override(self) -> None: try: body = self.read_json_body() name = validate_text(body.get("name"), "行业或题材名称", 50, required=True) element = str(body.get("element") or "").strip() if element not in {"木", "火", "土", "金", "水"}: raise ValueError("五行归类必须是木、火、土、金或水。") SERVICE.database.save_sector_phase_override(name, element) self.send_json({"ok": True}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def backfill_data(self) -> None: try: body = self.read_json_body() results = SERVICE.backfill( str(body.get("start_date") or ""), str(body.get("end_date") or ""), ) self.send_json({"ok": True, "results": results}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) except Exception as exc: self.send_json({"error": f"历史回补失败:{exc}"}, HTTPStatus.INTERNAL_SERVER_ERROR) def sync_screener_data(self) -> None: try: body = self.read_json_body() result = SERVICE.sync_screener_data( str(body.get("trade_date") or date.today().isoformat()), int(body.get("lookback") or 45), ) self.send_json({"ok": True, "result": result}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) except Exception as exc: self.send_json({"error": f"因子数据同步失败:{exc}"}, HTTPStatus.INTERNAL_SERVER_ERROR) def compile_screener_strategy(self) -> None: try: body = self.read_json_body() result = SERVICE.compile_screener_strategy( str(body.get("prompt") or ""), str(body.get("regime") or "") ) self.send_json({"ok": True, "strategy": result}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def save_screener_strategy(self) -> None: try: body = self.read_json_body() result = SERVICE.save_screener_strategy(body) self.send_json({"ok": True, **result}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def run_screener(self) -> None: try: body = self.read_json_body() result = SERVICE.run_screener(body) self.send_json({"ok": True, "result": result}) except ValueError as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) except Exception as exc: self.send_json({"error": f"选股执行失败:{exc}"}, HTTPStatus.INTERNAL_SERVER_ERROR) def mentor_chat(self) -> None: try: body = self.read_json_body() result = SERVICE.mentor_chat(body) self.send_json({"ok": True, **result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def heaven_hexagram(self) -> None: try: body = self.read_json_body() result = SERVICE.heaven_hexagram(body.get("lines")) self.send_json({"ok": True, "hexagram": result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def heaven_personal(self) -> None: try: body = self.read_json_body() result = SERVICE.heaven_personal(body) self.send_json({"ok": True, "personal": result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def heaven_interpret(self) -> None: try: body = self.read_json_body() result = SERVICE.heaven_interpret(body) self.send_json({"ok": True, **result}) except (ValueError, json.JSONDecodeError) as exc: self.send_json({"error": str(exc)}, HTTPStatus.BAD_REQUEST) def read_json_body(self, allow_empty: bool = False) -> dict[str, Any]: length = int(self.headers.get("Content-Length", "0")) if length == 0 and allow_empty: return {} if length <= 0 or length > 65536: raise ValueError("请求内容为空或过大。") return json.loads(self.rfile.read(length).decode("utf-8")) def serve_static(self, request_path: str) -> None: relative = unquote(request_path).lstrip("/") or "index.html" candidate = (STATIC_DIR / relative).resolve() try: candidate.relative_to(STATIC_DIR.resolve()) except ValueError: self.send_error(HTTPStatus.FORBIDDEN) return if not candidate.is_file(): candidate = STATIC_DIR / "index.html" try: content = candidate.read_bytes() except OSError: self.send_error(HTTPStatus.NOT_FOUND) return content_type = mimetypes.guess_type(candidate.name)[0] or "application/octet-stream" if content_type.startswith("text/") or content_type in {"application/javascript", "application/json"}: content_type += "; charset=utf-8" self.send_response(HTTPStatus.OK) self.send_header("Content-Type", content_type) self.send_header("Content-Length", str(len(content))) self.send_header("Cache-Control", "no-cache") self.end_headers() self.wfile.write(content) def send_json( self, payload: dict[str, Any], status: HTTPStatus = HTTPStatus.OK, headers: dict[str, str] | None = None, ) -> None: content = json.dumps(payload, ensure_ascii=False).encode("utf-8") self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(content))) self.send_header("Cache-Control", "no-store") for name, value in (headers or {}).items(): self.send_header(name, value) self.end_headers() self.wfile.write(content) def log_message(self, format_string: str, *args: Any) -> None: print(f"[{self.log_date_time_string()}] {format_string % args}") def normalize_date(value: str) -> str: compact = value.replace("-", "").strip() try: parsed = datetime.strptime(compact, "%Y%m%d") except ValueError as exc: raise ValueError("日期格式应为 YYYY-MM-DD。") from exc if parsed.date() > date.today(): raise ValueError("不能查询未来日期。") return parsed.strftime("%Y%m%d") def validate_stock_code(value: str) -> str: code = value.strip() if not re.fullmatch(r"\d{6}", code): raise ValueError("股票代码应为 6 位数字。") return code def tushare_code(code: str) -> str: if code.startswith(("4", "8", "9")): suffix = "BJ" elif code.startswith("6"): suffix = "SH" else: suffix = "SZ" return f"{code}.{suffix}" def validate_text(value: Any, label: str, maximum: int, required: bool = False) -> str: text = str(value or "").strip() if required and not text: raise ValueError(f"{label}不能为空。") if len(text) > maximum: raise ValueError(f"{label}不能超过 {maximum} 个字符。") return text def _parse_iso_datetime(value: Any) -> datetime | None: text = str(value or "").strip() if not text: return None try: parsed = datetime.fromisoformat(text) except ValueError: return None return parsed.replace(tzinfo=timezone.utc) if parsed.tzinfo is None else parsed.astimezone(timezone.utc) def _membership_boundary(value: Any, end: bool) -> str | None: text = str(value or "").strip() if not text: return None try: day = datetime.strptime(text, "%Y-%m-%d").replace(tzinfo=timezone.utc) except ValueError as exc: raise ValueError("会员日期格式应为 YYYY-MM-DD。") from exc if end: day += timedelta(days=1) return day.isoformat(timespec="seconds") def _add_months(value: datetime, months: int) -> datetime: month_index = value.year * 12 + value.month - 1 + months year, zero_based_month = divmod(month_index, 12) month = zero_based_month + 1 day = min(value.day, calendar.monthrange(year, month)[1]) return value.replace(year=year, month=month, day=day) def main() -> None: parser = argparse.ArgumentParser(description="Xiaobai stock review web application") parser.add_argument("--host", default="127.0.0.1") parser.add_argument("--port", type=int, default=8765) args = parser.parse_args() server = ThreadingHTTPServer((args.host, args.port), RequestHandler) print(f"Xiaobai Review Web is running at http://{args.host}:{args.port}") print("Press Ctrl+C to stop.") try: server.serve_forever() except KeyboardInterrupt: pass finally: SERVICE._background_stop.set() server.server_close() if __name__ == "__main__": main()