项目目录从 refer 迁移到 local

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-07-04 16:32:48 +08:00
parent 648a8b7f1c
commit aad34202f1
302 changed files with 0 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""业务服务层。"""
+429
View File
@@ -0,0 +1,429 @@
"""AI provider adapter for OpenAI-compatible APIs and local Codex CLI."""
from __future__ import annotations
import asyncio
import os
import re
import shutil
import sys
import tempfile
import tomllib
from collections.abc import AsyncIterator, Sequence
from pathlib import Path
from app import secrets_store
from app.config import settings
OPENAI_COMPAT_PROVIDER = "openai_compat"
CODEX_CLI_PROVIDER = "codex_cli"
CODEX_DEFAULT_COMMAND = "codex"
CODEX_SERVICE_TIER_FALLBACK = "fast"
CODEX_SUPPORTED_SERVICE_TIERS = {"fast", "flex"}
Message = dict[str, str]
_ANSI_RE = re.compile(r"\x1b\[[0-9;?]*[ -/]*[@-~]")
def current_ai_provider() -> str:
return secrets_store.get_ai_config("ai_provider", settings.ai_provider) or OPENAI_COMPAT_PROVIDER
def current_ai_model() -> str:
if current_ai_provider() == CODEX_CLI_PROVIDER:
return normalize_codex_model(str(secrets_store.load().get("ai_model") or ""))
return secrets_store.get_ai_config("ai_model", settings.ai_model)
def current_codex_command() -> str:
return normalize_codex_command(
secrets_store.get_ai_config("ai_codex_command", settings.ai_codex_command),
strict=False,
)
def is_codex_cli_provider(provider: str | None = None) -> bool:
return (provider or current_ai_provider()) == CODEX_CLI_PROVIDER
def normalize_codex_model(model: str) -> str:
value = model.strip()
aliases = {
"gpt5": "gpt-5",
"gpt5.5": "gpt-5.5",
}
return aliases.get(value.lower(), value)
def normalize_codex_command(command: str | None, *, strict: bool = True) -> str:
value = (command or "").strip()
if not value or value.lower() == CODEX_DEFAULT_COMMAND:
return CODEX_DEFAULT_COMMAND
if strict:
raise ValueError("Codex CLI 仅支持使用默认 codex 命令自动解析, 不支持自定义可执行路径")
return CODEX_DEFAULT_COMMAND
def normalize_openai_base_url(url: str) -> str:
"""Return the OpenAI-compatible base URL expected by the OpenAI SDK."""
base = (url or "").strip().rstrip("/")
if base.endswith("/chat/completions"):
base = base[: -len("/chat/completions")].rstrip("/")
if not base.endswith("/v1"):
base = f"{base}/v1"
return base
def codex_cli_available() -> bool:
try:
_codex_base_command()
return True
except RuntimeError:
return False
def ai_configured(provider: str | None = None) -> bool:
provider = provider or current_ai_provider()
if is_codex_cli_provider(provider):
return codex_cli_available()
return bool(secrets_store.get_ai_key())
async def generate_ai_text(
messages: Sequence[Message],
*,
temperature: float = 0.3,
max_tokens: int = 3000,
timeout: float = 180.0,
) -> str:
"""Return a complete AI response from the currently configured provider."""
if is_codex_cli_provider():
return await _run_codex_cli(messages, max_tokens=max_tokens, timeout=max(timeout, 600.0))
return await _run_openai_once(
messages,
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
)
async def stream_ai_text(
messages: Sequence[Message],
*,
temperature: float = 0.5,
max_tokens: int = 4000,
timeout: float = 180.0,
) -> AsyncIterator[str]:
"""Yield text deltas from the configured provider.
Codex CLI only exposes the final assistant message for this use case, so it
yields one complete chunk after the command exits.
"""
if is_codex_cli_provider():
yield await _run_codex_cli(messages, max_tokens=max_tokens, timeout=max(timeout, 600.0))
return
async for chunk in _stream_openai(
messages,
temperature=temperature,
max_tokens=max_tokens,
timeout=timeout,
):
yield chunk
async def _run_openai_once(
messages: Sequence[Message],
*,
temperature: float,
max_tokens: int,
timeout: float,
) -> str:
ai_key = secrets_store.get_ai_key()
if not ai_key:
raise RuntimeError("AI API Key 未配置, 请在设置页配置")
client = _openai_client(ai_key, timeout)
resp = await client.chat.completions.create(
model=current_ai_model(),
messages=list(messages),
temperature=temperature,
max_tokens=max_tokens,
)
if not resp.choices:
return ""
return (resp.choices[0].message.content or "").strip()
async def _stream_openai(
messages: Sequence[Message],
*,
temperature: float,
max_tokens: int,
timeout: float,
) -> AsyncIterator[str]:
ai_key = secrets_store.get_ai_key()
if not ai_key:
raise RuntimeError("AI API Key 未配置, 请在设置页配置")
client = _openai_client(ai_key, timeout)
stream = await client.chat.completions.create(
model=current_ai_model(),
messages=list(messages),
temperature=temperature,
max_tokens=max_tokens,
stream=True,
)
async for chunk in stream:
delta = chunk.choices[0].delta if chunk.choices else None
if delta and delta.content:
yield delta.content
def _openai_client(api_key: str, timeout: float):
from openai import AsyncOpenAI
user_agent = secrets_store.get_ai_config("ai_user_agent", "") or settings.ai_user_agent
return AsyncOpenAI(
api_key=api_key,
base_url=normalize_openai_base_url(secrets_store.get_ai_config("ai_base_url", settings.ai_base_url)),
timeout=timeout,
max_retries=2,
default_headers={"User-Agent": user_agent},
)
async def _run_codex_cli(
messages: Sequence[Message],
*,
max_tokens: int,
timeout: float,
) -> str:
prompt = _codex_prompt(messages, max_tokens=max_tokens)
with tempfile.TemporaryDirectory(prefix="tickflow-codex-run-") as run_dir:
run_path = Path(run_dir)
codex_home_path = run_path / "codex-home"
workspace_path = run_path / "workspace"
codex_home_path.mkdir()
workspace_path.mkdir()
output_path = codex_home_path / "last-message.txt"
_prepare_codex_home(codex_home_path)
args = [
*_codex_base_command(),
"exec",
"--ephemeral",
"--sandbox",
"read-only",
"--skip-git-repo-check",
"--color",
"never",
"--output-last-message",
str(output_path),
]
model = current_ai_model().strip()
if model:
args.extend(["--model", model])
args.extend(["--cd", str(workspace_path), "-"])
env = os.environ.copy()
env.setdefault("NO_COLOR", "1")
env["CODEX_HOME"] = str(codex_home_path)
proc = await asyncio.create_subprocess_exec(
*args,
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
env=env,
)
try:
stdout, stderr = await asyncio.wait_for(
proc.communicate(prompt.encode("utf-8")),
timeout=timeout,
)
except TimeoutError as exc:
proc.kill()
await proc.wait()
raise RuntimeError("Codex CLI 调用超时, 请稍后重试或检查本机 Codex 登录状态") from exc
out = _clean_process_text(stdout)
err = _clean_process_text(stderr)
final_message = _read_output_file(output_path)
if proc.returncode != 0:
detail = err or out or f"exit code {proc.returncode}"
raise RuntimeError(f"Codex CLI 调用失败: {detail[-1200:]}")
result = final_message or out
if not result:
raise RuntimeError("Codex CLI 未返回内容")
return result
def _codex_prompt(messages: Sequence[Message], *, max_tokens: int) -> str:
parts = [
"You are TickFlow Stock Panel's local AI provider.",
"This is a text-generation task. The working directory is intentionally empty.",
"Use only the user-provided prompt content below; do not inspect or modify local files.",
"Return only the final requested content; do not include execution logs.",
]
if max_tokens > 0:
parts.append(f"Keep the final answer within about {max_tokens} output tokens.")
for message in messages:
role = message.get("role", "user")
content = message.get("content", "")
parts.append(f"\n<{role}>\n{content}\n</{role}>")
return "\n".join(parts)
def _codex_base_command() -> list[str]:
command = current_codex_command()
resolved = _resolve_command(command)
if not resolved:
raise RuntimeError(f"未找到 Codex CLI 命令: {command}")
if sys.platform == "win32" and resolved.lower().endswith(".ps1"):
return ["powershell.exe", "-NoProfile", "-ExecutionPolicy", "Bypass", "-File", resolved]
return [resolved]
def _resolve_command(command: str) -> str | None:
if command.lower() != CODEX_DEFAULT_COMMAND:
return None
if sys.platform == "win32":
desktop_codex = _resolve_windows_desktop_codex()
if desktop_codex:
return desktop_codex
resolved = shutil.which(command)
if sys.platform == "win32" and resolved:
resolved_path = Path(resolved)
if not resolved_path.suffix:
cmd_path = resolved_path.with_suffix(".cmd")
if cmd_path.exists():
return str(cmd_path)
if not resolved and sys.platform == "win32" and not command.lower().endswith(".cmd"):
resolved = shutil.which(f"{command}.cmd")
if not resolved and sys.platform == "win32":
resolved = _resolve_windows_codex_command(command)
return resolved
def _resolve_windows_codex_command(command: str) -> str | None:
"""Find npm-installed Codex when the backend process has a minimal PATH."""
raw = Path(command)
if raw.parent != Path("."):
return None
names = [command]
if not raw.suffix:
names = [f"{command}.cmd", f"{command}.exe", f"{command}.bat", f"{command}.ps1", command]
dirs: list[Path] = []
appdata = os.environ.get("APPDATA")
if appdata:
dirs.append(Path(appdata) / "npm")
dirs.append(Path.home() / "AppData" / "Roaming" / "npm")
for env_name in ("ProgramFiles", "ProgramFiles(x86)", "LOCALAPPDATA"):
value = os.environ.get(env_name)
if value:
dirs.append(Path(value) / "nodejs")
for directory in dirs:
for name in names:
candidate = directory / name
if candidate.exists():
return str(candidate)
return None
def _resolve_windows_desktop_codex() -> str | None:
"""Prefer the Codex Desktop bundled CLI over an older npm shim."""
local_appdata = os.environ.get("LOCALAPPDATA")
if not local_appdata:
return None
root = Path(local_appdata) / "OpenAI" / "Codex" / "bin"
if not root.exists():
return None
candidates = list(root.glob("*/codex.exe"))
direct = root / "codex.exe"
if direct.exists():
candidates.append(direct)
if not candidates:
return None
newest = max(candidates, key=lambda p: p.stat().st_mtime)
return str(newest)
def _prepare_codex_home(target: Path) -> None:
"""Create an isolated CODEX_HOME that reuses auth but not fragile config."""
source = _codex_home()
auth_file = source / "auth.json"
if auth_file.exists():
shutil.copy2(auth_file, target / "auth.json")
_write_compatible_codex_config(target / "config.toml")
def _codex_home() -> Path:
return Path(os.environ.get("CODEX_HOME") or Path.home() / ".codex")
def _write_compatible_codex_config(path: Path) -> None:
config = _read_codex_config()
lines: list[str] = []
tier = str(config.get("service_tier") or "").strip()
if tier not in CODEX_SUPPORTED_SERVICE_TIERS:
tier = CODEX_SERVICE_TIER_FALLBACK
lines.append(_toml_string("service_tier", tier))
lines.append(_toml_string("approval_policy", "never"))
lines.append(_toml_string("sandbox_mode", "read-only"))
path.write_text("\n".join(lines) + "\n", encoding="utf-8")
def _read_codex_config() -> dict:
path = _codex_home() / "config.toml"
if not path.exists():
return {}
try:
with path.open("rb") as f:
return tomllib.load(f)
except tomllib.TOMLDecodeError:
return _read_codex_config_lenient(path)
except OSError:
return {}
def _read_codex_config_lenient(path: Path) -> dict:
config: dict[str, str] = {}
pattern = re.compile(r'^\s*([A-Za-z0-9_-]+)\s*=\s*"([^"]*)"\s*$')
try:
for line in path.read_text(encoding="utf-8", errors="replace").splitlines():
match = pattern.match(line)
if match:
config[match.group(1)] = match.group(2)
except OSError:
pass
return config
def _toml_string(key: str, value: str) -> str:
escaped = value.replace("\\", "\\\\").replace('"', '\\"')
return f'{key} = "{escaped}"'
def _clean_process_text(raw: bytes) -> str:
text = raw.decode("utf-8", errors="replace")
return _ANSI_RE.sub("", text).strip()
def _read_output_file(path: Path) -> str:
if path.exists():
return _ANSI_RE.sub("", path.read_text(encoding="utf-8", errors="replace")).strip()
return ""
+101
View File
@@ -0,0 +1,101 @@
"""AI 财务分析报告持久化存储。
存储位置: data/user_data/ai_reports.json (数组,按 created_at 降序)
保留最近 MAX_REPORTS 条;超出自动裁剪最旧的。
每条报告结构:
{
"id": "rpt_xxx", # 唯一 id
"symbol": "600519.SH",
"name": "贵州茅台",
"focus": "", # 用户追加的关心点(可为空)
"content": "# ...markdown", # 报告正文
"periods": 4, # 基于几期数据生成
"summary": "metrics: 1期...", # 数据摘要
"created_at": "2026-06-25T10:00:00"
}
"""
from __future__ import annotations
import json
import logging
import time
from pathlib import Path
logger = logging.getLogger(__name__)
MAX_REPORTS = 20
def _path() -> Path:
from app.config import settings
p = settings.data_dir / "user_data" / "ai_reports.json"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def list_reports() -> list[dict]:
"""返回全部报告(按 created_at 降序)。"""
p = _path()
if not p.exists():
return []
try:
data = json.loads(p.read_text(encoding="utf-8"))
if isinstance(data, list):
return sorted(data, key=lambda r: r.get("created_at", ""), reverse=True)
except Exception as e: # noqa: BLE001
logger.warning("ai_reports.json malformed: %s", e)
return []
def _save_all(reports: list[dict]) -> None:
"""全量写入(裁剪到 MAX_REPORTS)。"""
# 保持降序
reports.sort(key=lambda r: r.get("created_at", ""), reverse=True)
if len(reports) > MAX_REPORTS:
reports = reports[:MAX_REPORTS]
_path().write_text(
json.dumps(reports, indent=2, ensure_ascii=False), encoding="utf-8",
)
def save_report(report: dict) -> dict:
"""新增一条报告并持久化。返回保存后的报告(含 id / created_at)。
自动补全 id 与 created_at(若缺),并裁剪到上限。
"""
reports = list_reports()
if not report.get("id"):
report["id"] = f"rpt_{int(time.time() * 1000)}_{report.get('symbol', 'x')}"
if not report.get("created_at"):
report["created_at"] = _now_iso()
reports.append(report)
_save_all(reports)
logger.info("AI report saved: %s (%s), total %d", report.get("symbol"), report.get("id"), len(reports))
return report
def delete_report(report_id: str) -> bool:
"""删除指定报告。返回是否删除成功。"""
reports = list_reports()
before = len(reports)
reports = [r for r in reports if r.get("id") != report_id]
if len(reports) < before:
_save_all(reports)
return True
return False
def clear_reports() -> int:
"""清空全部报告。返回删除数量。"""
reports = list_reports()
n = len(reports)
if n > 0:
_save_all([])
return n
def _now_iso() -> str:
"""当前本地时间 ISO 字符串(带秒精度,前端 toLocaleString 友好)。"""
from datetime import datetime
return datetime.now().isoformat(timespec="seconds")
+209
View File
@@ -0,0 +1,209 @@
"""告警触发记录存储 — JSONL 追加写 + 滚动清理。
职责:
- 把每次触发的 AlertEvent 追加写入 data/user_data/alerts.jsonl
- 提供查询 (按来源/类型过滤、时间倒序、限量)
- 滚动清理: 保留近 N 天 + 上限 M 条 (取交集)
设计:
- JSONL 每行一个 JSON 对象,便于增量追加和流式读取
- 清理策略: 追加后按需 prune (按 ts 删旧),避免文件无限膨胀
- 读时全量加载到内存过滤 (记录量受上限约束, 5000 条量级无压力)
"""
from __future__ import annotations
import json
import logging
import threading
from pathlib import Path
logger = logging.getLogger(__name__)
# 保留策略
MAX_DAYS = 7
MAX_RECORDS = 5000
# 每隔多少次写入触发一次清理 (避免每次写都 prune)
PRUNE_EVERY = 20
_lock = threading.Lock()
_write_count = 0
def _path(data_dir: Path) -> Path:
p = data_dir / "user_data" / "alerts.jsonl"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def append(data_dir: Path, event: dict) -> None:
"""追加一条触发记录。event 应含 ts(毫秒)、rule_id、source 等字段。"""
line = json.dumps(event, ensure_ascii=False)
with _lock:
p = _path(data_dir)
with p.open("a", encoding="utf-8") as f:
f.write(line + "\n")
global _write_count
_write_count += 1
if _write_count >= PRUNE_EVERY:
_write_count = 0
_prune_locked(p)
def append_many(data_dir: Path, events: list[dict]) -> None:
"""批量追加。"""
if not events:
return
with _lock:
p = _path(data_dir)
with p.open("a", encoding="utf-8") as f:
for ev in events:
f.write(json.dumps(ev, ensure_ascii=False) + "\n")
global _write_count
_write_count += len(events)
if _write_count >= PRUNE_EVERY:
_write_count = 0
_prune_locked(p)
def list_recent(
data_dir: Path,
days: int = MAX_DAYS,
limit: int = MAX_RECORDS,
source: str | None = None,
type: str | None = None,
) -> list[dict]:
"""读取近 N 天记录,按时间倒序,支持按 source/type 过滤。"""
import time
cutoff = (time.time() - days * 86400) * 1000 # 毫秒
out: list[dict] = []
p = _path(data_dir)
if not p.exists():
return []
try:
with p.open("r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
ev = json.loads(line)
except Exception:
continue
if ev.get("ts", 0) < cutoff:
continue
if source and ev.get("source") != source:
continue
if type and ev.get("type") != type:
continue
out.append(ev)
except Exception as e:
logger.warning("alert_store read failed: %s", e)
return []
# 时间倒序 + 截断
out.sort(key=lambda x: x.get("ts", 0), reverse=True)
return out[:limit]
def clear(data_dir: Path) -> int:
"""清空全部记录,返回清除的条数。"""
with _lock:
p = _path(data_dir)
if not p.exists():
return 0
count = 0
try:
with p.open("r", encoding="utf-8") as f:
count = sum(1 for line in f if line.strip())
except Exception:
pass
p.write_text("", encoding="utf-8")
return count
def delete_one(data_dir: Path, ts: int) -> bool:
"""删除指定 ts 的单条记录,返回是否删除成功。
JSONL 无主键, 用 ts(毫秒时间戳) 作为标识。
若存在多条同 ts, 只删第一条。
"""
with _lock:
p = _path(data_dir)
if not p.exists():
return False
kept: list[dict] = []
deleted = False
try:
with p.open("r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
ev = json.loads(line)
except Exception:
continue
if not deleted and ev.get("ts") == ts:
deleted = True
continue
kept.append(ev)
except Exception as e:
logger.warning("alert_store delete_one read failed: %s", e)
return False
if not deleted:
return False
try:
with p.open("w", encoding="utf-8") as f:
for ev in kept:
f.write(json.dumps(ev, ensure_ascii=False) + "\n")
except Exception as e:
logger.warning("alert_store delete_one write failed: %s", e)
return False
return True
return count
def count(data_dir: Path) -> int:
"""返回当前记录总数。"""
p = _path(data_dir)
if not p.exists():
return 0
try:
with p.open("r", encoding="utf-8") as f:
return sum(1 for line in f if line.strip())
except Exception:
return 0
def _prune_locked(p: Path) -> None:
"""(调用方需持锁) 保留近 MAX_DAYS 天 + 上限 MAX_RECORDS 条。"""
import time
cutoff = (time.time() - MAX_DAYS * 86400) * 1000
kept: list[dict] = []
try:
with p.open("r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
ev = json.loads(line)
except Exception:
continue
if ev.get("ts", 0) >= cutoff:
kept.append(ev)
except FileNotFoundError:
return
except Exception as e:
logger.warning("alert_store prune read failed: %s", e)
return
# 上限截断 (保留最新的)
if len(kept) > MAX_RECORDS:
kept.sort(key=lambda x: x.get("ts", 0))
kept = kept[-MAX_RECORDS:]
# 重写文件
try:
with p.open("w", encoding="utf-8") as f:
for ev in kept:
f.write(json.dumps(ev, ensure_ascii=False) + "\n")
except Exception as e:
logger.warning("alert_store prune write failed: %s", e)
+201
View File
@@ -0,0 +1,201 @@
"""访问密码认证 — 单用户, 自托管场景。
设计:
- 密码用 PBKDF2-HMAC-SHA256 哈希(标准库 hashlib, 无新依赖), 加随机 salt。
即使 auth.json 泄露, 也无法逆向出明文密码。
- 会话用随机 token(token_urlsafe), 内存 + 文件双存(支持多进程/重启不丢失)。
- 存储: data/user_data/auth.json (chmod 0600), 仿 secrets_store 模式。
安全要点:
- 设密码接口必须限制本机/内网(见 auth router), 防黑客抢占域名抢先设密码。
- 登录限流: 错5次锁5分钟(见 auth router 内存计数)。
- 单密码, 不做多用户(避免重构全项目数据层)。
"""
from __future__ import annotations
import hashlib
import json
import logging
import os
import secrets as _secrets
import threading
import time
from pathlib import Path
logger = logging.getLogger(__name__)
# PBKDF2 参数(NIST 推荐, 单次校验 ~100ms, 兼顾安全与响应)
_PBKDF2_ITER = 200_000
_SALT_LEN = 16
_TOKEN_BYTES = 32
# 会话有效期: 30 天(自托管单用户, 长一点减少重登频率)
SESSION_TTL = 30 * 24 * 3600
_lock = threading.Lock()
# 内存中的有效会话: { token: expire_ts }。进程重启后从磁盘恢复。
_sessions: dict[str, float] = {}
def _path() -> Path:
from app.config import settings
p = settings.data_dir / "user_data" / "auth.json"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def _load() -> dict:
p = _path()
if p.exists():
try:
return json.loads(p.read_text(encoding="utf-8"))
except Exception as e: # noqa: BLE001
logger.warning("auth.json malformed: %s", e)
return {}
def _save(data: dict) -> None:
p = _path()
p.write_text(json.dumps(data, indent=2, ensure_ascii=False), encoding="utf-8")
try:
os.chmod(p, 0o600)
except OSError:
pass
def _hash_password(password: str, salt: bytes | None = None) -> tuple[str, str]:
"""返回 (salt_hex, hash_hex)。salt 为 None 时生成新 salt。"""
if salt is None:
salt = os.urandom(_SALT_LEN)
dk = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, _PBKDF2_ITER)
return salt.hex(), dk.hex()
def _verify_password(password: str, salt_hex: str, hash_hex: str) -> bool:
"""恒定时间比较, 防时序攻击。"""
try:
salt = bytes.fromhex(salt_hex)
expected = bytes.fromhex(hash_hex)
except ValueError:
return False
actual = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, _PBKDF2_ITER)
return _secrets.compare_digest(actual, expected)
# ================================================================
# 密码管理
# ================================================================
def is_configured() -> bool:
"""是否已设置访问密码。"""
d = _load()
return bool(d.get("password_hash"))
def set_password(password: str) -> None:
"""设置/修改访问密码。清空所有现有会话(强制重新登录)。"""
if len(password) < 6:
raise ValueError("密码至少 6 位")
salt_hex, hash_hex = _hash_password(password)
with _lock:
_sessions.clear() # 改密码 = 旧会话全部失效
_save({
"password_hash": hash_hex,
"password_salt": salt_hex,
"updated_at": int(time.time()),
"sessions": {}, # 清空持久化会话
})
logger.info("access password set")
def bootstrap_from_env() -> bool:
"""首次初始化: 若环境变量 AUTH_PASSWORD 已配置且尚未设过密码, 则用它设密码。
公网服务器部署场景: 避免每次都要 SSH 端口转发才能设首个密码。
明文密码只在内存/配置中, 经 set_password() 哈希后写入 auth.json (chmod 0600)。
一旦设置成功, 后续重启不再覆盖 (用户改密码走 UI, 不受环境变量影响)。
Returns:
True 表示本次用环境变量初始化了密码; False 表示无需初始化。
"""
from app.config import settings
pwd = (settings.auth_password or "").strip()
if not pwd:
return False
if is_configured():
# 已设过密码, 不覆盖 (避免环境变量反复重置用户在 UI 改的密码)
return False
try:
set_password(pwd)
logger.info("access password bootstrapped from AUTH_PASSWORD env (one-time)")
return True
except ValueError as e:
# 密码不合规 (< 6 位), 记日志但不阻断启动
logger.warning("AUTH_PASSWORD bootstrap skipped: %s", e)
return False
def verify_and_create_session(password: str) -> str | None:
"""验证密码, 成功则创建会话并返回 token, 失败返回 None。"""
d = _load()
if not d.get("password_hash"):
return None
if not _verify_password(password, d.get("password_salt", ""), d["password_hash"]):
return None
token = _secrets.token_urlsafe(_TOKEN_BYTES)
expire = time.time() + SESSION_TTL
with _lock:
_sessions[token] = expire
_persist_sessions_locked()
return token
def revoke_session(token: str) -> None:
"""注销会话(登出)。"""
with _lock:
_sessions.pop(token, None)
_persist_sessions_locked()
def is_valid_session(token: str) -> bool:
"""检查会话是否有效(存在且未过期)。过期则清理。"""
if not token:
return False
with _lock:
expire = _sessions.get(token)
if expire is None:
return False
if time.time() > expire:
_sessions.pop(token, None)
_persist_sessions_locked()
return False
return True
def _persist_sessions_locked() -> None:
"""把当前内存会话写回 auth.json(需持锁调用)。"""
d = _load()
d["sessions"] = {t: exp for t, exp in _sessions.items()}
_save(d)
def _restore_sessions() -> None:
"""启动时从 auth.json 恢复未过期会话(支持进程重启不丢登录态)。"""
with _lock:
d = _load()
now = time.time()
saved = d.get("sessions") or {}
for token, expire in saved.items():
if isinstance(expire, (int, float)) and expire > now:
_sessions[token] = expire
if len(_sessions) != len(saved):
# 有过期会话被清理, 落盘一次
_persist_sessions_locked()
# 模块加载时恢复会话
try:
_restore_sessions()
except Exception as e: # noqa: BLE001
logger.warning("restore sessions failed: %s", e)
+392
View File
@@ -0,0 +1,392 @@
"""回测服务(§6.7)。
包 vectorbt — 全项目唯一一处出现 pandas。
"""
from __future__ import annotations
import logging
import uuid
from dataclasses import dataclass, field
from datetime import date
from typing import Literal
import numpy as np
import pandas as pd
import polars as pl
from app.config import settings
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
# vectorbt 是 optional extras(见 pyproject.toml).未装时只有 backtest 不可用,其他功能正常.
_vbt = None
_vbt_unavailable_reason: str | None = None
class VectorbtUnavailable(RuntimeError):
"""vectorbt 未安装 — 提示用户 `uv sync --extra backtest`."""
def _get_vbt():
global _vbt, _vbt_unavailable_reason
if _vbt is not None:
return _vbt
if _vbt_unavailable_reason is not None:
raise VectorbtUnavailable(_vbt_unavailable_reason)
try:
import vectorbt as vbt
_vbt = vbt
return _vbt
except ImportError as e:
_vbt_unavailable_reason = (
"vectorbt 未安装 — 它是回测的可选依赖.macOS Intel 用户先 `brew install cmake` "
"然后 `uv sync --extra backtest`"
)
logger.warning("vectorbt unavailable: %s", e)
raise VectorbtUnavailable(_vbt_unavailable_reason) from e
def is_available() -> bool:
"""供 API 层快速检测."""
try:
_get_vbt()
return True
except VectorbtUnavailable:
return False
SignalKind = Literal[
"macd_golden", "macd_dead",
"ma_golden_5_20", "ma_dead_5_20",
"ma_golden_20_60",
"ma20_breakout", "ma20_breakdown",
"n_day_high", "n_day_low",
"boll_breakout_upper", "boll_breakdown_lower",
"volume_surge",
"rsi_oversold", "rsi_overbought",
"stop_loss", "trailing_stop", "max_hold",
]
@dataclass
class BacktestConfig:
symbols: list[str]
start: date
end: date
# 买入信号(任一触发即买)
entries: list[str] = field(default_factory=list)
# 卖出信号(任一触发即卖)
exits: list[str] = field(default_factory=list)
# 其他参数
stop_loss_pct: float | None = None # 例 -0.05 = -5%
max_hold_days: int | None = None
fees_pct: float = 0.0002 # 万二佣金
slippage_bps: float = 5 # 5 bps
# 撮合
matching: Literal["close_t", "open_t+1"] = "close_t"
rsi_oversold_threshold: float = 30
rsi_overbought_threshold: float = 70
@dataclass
class BacktestResult:
run_id: str
config: dict
stats: dict
equity_curve: list[dict] # [{date, value}]
trades: list[dict] # [{symbol, entry_date, exit_date, pnl_pct, ...}]
per_symbol_stats: list[dict] # 每只股票的统计
# enriched 表里的信号列名映射
_SIGNAL_COLS: dict[SignalKind, str] = {
"macd_golden": "signal_macd_golden",
"macd_dead": "signal_macd_dead",
"ma_golden_5_20": "signal_ma_golden_5_20",
"ma_dead_5_20": "signal_ma_dead_5_20",
"ma_golden_20_60": "signal_ma_golden_20_60",
"ma20_breakout": "signal_ma20_breakout",
"ma20_breakdown": "signal_ma20_breakdown",
"n_day_high": "signal_n_day_high",
"n_day_low": "signal_n_day_low",
"boll_breakout_upper": "signal_boll_breakout_upper",
"boll_breakdown_lower": "signal_boll_breakdown_lower",
"volume_surge": "signal_volume_surge",
}
class BacktestService:
def __init__(self, repo: KlineRepository) -> None:
self.repo = repo
def _load_panel(
self,
symbols: list[str],
start: date,
end: date,
) -> pd.DataFrame:
"""加载 [date × symbol] 价格面板 — Polars scan_parquet + 即时计算指标。
**全项目唯一从 Polars 转 pandas 的边界**(§7.4 / ADR-19)。
"""
try:
enriched_glob = str(self.repo.store.data_dir / "kline_daily_enriched" / "**" / "*.parquet")
df = (
pl.scan_parquet(enriched_glob)
.filter(
(pl.col("symbol").is_in(symbols))
& (pl.col("date") >= start)
& (pl.col("date") <= end)
)
.sort(["date", "symbol"])
.collect()
)
except Exception as e: # noqa: BLE001
logger.warning("backtest load failed: %s", e)
return pd.DataFrame()
if df.is_empty():
return pd.DataFrame()
# 即时计算指标 + 信号
from app.indicators.pipeline import compute_all
df = compute_all(df)
# 选择需要的列
needed_cols = [
"date", "symbol", "open", "high", "low", "close", "volume",
"rsi_14", "signal_macd_golden", "signal_macd_dead",
"signal_ma_golden_5_20", "signal_ma_dead_5_20",
"signal_ma_golden_20_60",
"signal_ma20_breakout", "signal_ma20_breakdown",
"signal_n_day_high", "signal_n_day_low",
"signal_boll_breakout_upper", "signal_boll_breakdown_lower",
"signal_volume_surge",
]
existing = [c for c in needed_cols if c in df.columns]
df = df.select(existing)
# to_pandas 边界
return df.to_pandas(use_pyarrow_extension_array=False)
def _build_signal_matrix(
self,
panel: pd.DataFrame,
kinds: list[str],
config: BacktestConfig,
) -> pd.DataFrame:
"""从面板构造 [date × symbol] 的布尔信号矩阵。"""
if not kinds or panel.empty:
return pd.DataFrame()
# pivot 成 [date × symbol] 形式
result = None
for kind in kinds:
mat = None
if kind in _SIGNAL_COLS:
col = _SIGNAL_COLS[kind]
mat = panel.pivot(index="date", columns="symbol", values=col).fillna(False).astype(bool)
elif kind == "rsi_oversold":
mat = (panel.pivot(index="date", columns="symbol", values="rsi_14")
< config.rsi_oversold_threshold)
elif kind == "rsi_overbought":
mat = (panel.pivot(index="date", columns="symbol", values="rsi_14")
> config.rsi_overbought_threshold)
# stop_loss / trailing / max_hold 通过 vectorbt 参数处理,不参与信号矩阵
if mat is not None:
result = mat if result is None else (result | mat)
return result if result is not None else pd.DataFrame()
def run(self, config: BacktestConfig) -> BacktestResult:
vbt = _get_vbt()
run_id = uuid.uuid4().hex[:10]
panel = self._load_panel(config.symbols, config.start, config.end)
if panel.empty:
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": "no data"},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# 价格面板
close = panel.pivot(index="date", columns="symbol", values="close")
# 信号矩阵
entries = self._build_signal_matrix(panel, config.entries, config)
exits = self._build_signal_matrix(panel, config.exits, config)
# 对齐 index/columns
if not entries.empty:
entries = entries.reindex_like(close).fillna(False).astype(bool)
else:
entries = pd.DataFrame(False, index=close.index, columns=close.columns)
if not exits.empty:
exits = exits.reindex_like(close).fillna(False).astype(bool)
else:
exits = pd.DataFrame(False, index=close.index, columns=close.columns)
if not entries.any().any():
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": "no buy signals"},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# T+1 适配:vectorbt 默认信号当根 K 撮合
# close_t 撮合:维持默认
# open_t+1 撮合:shift 信号 1 根 + 用 open 作为价
if config.matching == "open_t+1":
entries = entries.shift(1).fillna(False).astype(bool)
exits = exits.shift(1).fillna(False).astype(bool)
price = panel.pivot(index="date", columns="symbol", values="open")
else:
price = close
# 跑回测
try:
pf_kwargs = dict(
close=close,
entries=entries,
exits=exits,
price=price,
fees=config.fees_pct,
slippage=config.slippage_bps / 10000.0,
freq="1D",
)
if config.stop_loss_pct is not None:
pf_kwargs["sl_stop"] = abs(config.stop_loss_pct)
if config.max_hold_days is not None:
# vectorbt 没有内置 max-hold;用时间退出近似:
# 在 max_hold_days 后强制 exit
exits_idx = entries.copy()
for col in entries.columns:
entry_rows = np.where(entries[col].values)[0]
for i in entry_rows:
end_i = min(i + config.max_hold_days, len(entries) - 1)
if end_i > i:
exits_idx.iloc[end_i][col] = True
pf_kwargs["exits"] = (exits | exits_idx).astype(bool)
pf = vbt.Portfolio.from_signals(**pf_kwargs)
except Exception as e: # noqa: BLE001
logger.exception("vectorbt backtest failed")
return BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={"error": str(e)},
equity_curve=[],
trades=[],
per_symbol_stats=[],
)
# 提取结果
try:
stats_series = pf.stats(silence_warnings=True)
if isinstance(stats_series, pd.DataFrame):
# 多列时取 agg
stats_dict = stats_series.mean(numeric_only=True).to_dict()
else:
stats_dict = stats_series.to_dict()
except Exception: # noqa: BLE001
stats_dict = {}
# 净值曲线(组合平均)
equity = pf.value().mean(axis=1) if isinstance(pf.value(), pd.DataFrame) else pf.value()
equity_curve = [
{"date": str(idx.date() if hasattr(idx, "date") else idx), "value": float(v)}
for idx, v in equity.items() if pd.notna(v)
]
# 交易记录
try:
trades_df = pf.trades.records_readable
trades = trades_df.to_dict(orient="records") if not trades_df.empty else []
# 字段名美化
trades = [
{
"symbol": t.get("Column", t.get("Symbol", "")),
"entry_date": str(t.get("Entry Timestamp", t.get("Entry Date", ""))),
"exit_date": str(t.get("Exit Timestamp", t.get("Exit Date", ""))),
"entry_price": float(t.get("Avg Entry Price", t.get("Avg. Entry Price", 0))),
"exit_price": float(t.get("Avg Exit Price", t.get("Avg. Exit Price", 0))),
"pnl_pct": float(t.get("Return", t.get("PnL %", 0))),
"duration": str(t.get("Duration", "")),
}
for t in trades
]
except Exception: # noqa: BLE001
trades = []
# 每标的统计
per_symbol = []
try:
total_ret = pf.total_return()
if isinstance(total_ret, pd.Series):
for sym, ret in total_ret.items():
if pd.notna(ret):
per_symbol.append({"symbol": sym, "total_return": float(ret)})
except Exception: # noqa: BLE001
pass
result = BacktestResult(
run_id=run_id,
config=_config_to_dict(config),
stats={k: _json_safe(v) for k, v in stats_dict.items()},
equity_curve=equity_curve,
trades=trades,
per_symbol_stats=per_symbol,
)
# 落盘
self._persist(result)
return result
def _persist(self, result: BacktestResult) -> None:
out_dir = settings.data_dir / "backtest_results"
out_dir.mkdir(parents=True, exist_ok=True)
# 用 polars 写一份汇总
summary = pl.DataFrame({
"run_id": [result.run_id],
"stats_json": [str(result.stats)],
"n_trades": [len(result.trades)],
})
summary.write_parquet(out_dir / f"run_id={result.run_id}.parquet")
def get_result(self, run_id: str) -> BacktestResult | None:
# Phase 1:只保留近似落盘,完整结果保存在内存的近期 cache 中
# 简化:重新 run 比缓存复杂结果代价小,暂不实现 get_result
return None
def _config_to_dict(c: BacktestConfig) -> dict:
return {
"symbols": c.symbols,
"start": str(c.start),
"end": str(c.end),
"entries": c.entries,
"exits": c.exits,
"stop_loss_pct": c.stop_loss_pct,
"max_hold_days": c.max_hold_days,
"fees_pct": c.fees_pct,
"slippage_bps": c.slippage_bps,
"matching": c.matching,
}
def _json_safe(v):
if isinstance(v, (int, float, str, bool)) or v is None:
return v
if isinstance(v, (np.floating, np.integer)):
return float(v) if not np.isnan(float(v)) else None
if hasattr(v, "isoformat"):
return v.isoformat()
return str(v)
@@ -0,0 +1,357 @@
"""AI 概念轮动分析 — 从概念涨幅排名矩阵提炼主线/新晋/退潮信号。
数据来源:
- rps_rotation.build_rps_rotation: 概念涨幅排名矩阵 (N 日 × ~387 概念)
- market_overview_builder.build_market_overview: 大盘背景 (指数/情绪/涨停)
架构 (复刻 market_recap):
预计算轮动信号 → 拼装 prompt → stream_ai_text 流式调用 → NDJSON 协议输出
协议事件: meta(摘要) / delta(文本片段) / error / done
"""
from __future__ import annotations
import json
import logging
import math
from collections.abc import AsyncIterator
logger = logging.getLogger(__name__)
# ================================================================
# System Prompt — 轮动策略师人格 + 固定章节模板
# ================================================================
_SYSTEM_PROMPT = """你是一位专注 A 股题材轮动的资深策略师,拥有 12 年一线实战经验,擅长从概念板块的**涨幅排名矩阵**中识别主力资金脉络,区分机构主导的持续性主线与游资驱动的脉冲式轮动,产出可直接指导题材跟踪与节奏把握的轮动分析。
## 输出规范
用 **Markdown** 格式输出,严格遵循以下结构。不要输出任何 JSON 或代码块,直接输出 Markdown 正文。
### 1. 🎯 主线研判(2-3 句)
点名当前最核心的 1-2 条主线题材(连续多日霸榜的强势概念),用一句话概括其逻辑(政策/产业/业绩/事件驱动),并判断是**主升期/加速期/扩散期/见顶期**。结尾用【主线强度:强 / 中 / 弱】定性。
### 2. 🆕 新晋强势
列出排名快速跃升的概念(从榜单中后段冲进前列的),逐个给出:
- 概念名 + 近 N 日排名变化(如 `45→20→8`)
- 涨幅加速度(连日递增 = 趋势加强)
- 可能的驱动逻辑(从板块属性推断,不要编造具体消息)
- 判断是**主力切入**还是**消息脉冲**
### 3. 📉 退潮预警
列出从高位明显滑落的概念(连续排名下滑或涨幅骤降),逐个给出:
- 概念名 + 排名下滑轨迹
- 退潮性质(高位分歧/资金撤离/补跌)
- 是否扩散风险
### 4. 🏛️ 机构主线 vs 🎰 游资轮动
基于排名稳定性区分两类资金行为:
- **机构主线**:排名标准差小、长期稳居前列的概念 → 持续性判断、是否可作底仓方向
- **游资轮动**:排名剧烈波动、脉冲式冲高的概念 → 短线节奏提示、追高风险
给出当前市场**整体轮动节奏**(快轮动/慢轮动/主线聚焦)的判断。
### 5. 🌐 结合大盘
结合提供的大盘数据(指数涨跌/情绪/涨停数),判断:
- 当前大盘环境对题材轮动是助力还是阻力
- 情绪温度与轮动节奏的匹配度(如情绪冰点但题材活跃 = 抱团;情绪火热但轮动快 = 末段)
### 6. 🎯 操作建议
- **跟踪方向**:主线延续 + 新晋确认的概念
- **规避方向**:明确退潮 + 高位脉冲的概念
- **节奏提示**:当前适合追高 / 低吸 / 观望,及切换信号(如"主线概念连续 2 日跌出前 10 则确认退潮")
### 7. ⚠️ 风险提示
列出需要盯的风险(如主线断层、情绪与轮动背离、成交萎缩)。末尾附一行:
"> ⚠️ 本报告由 AI 基于公开行情数据生成,仅供参考,不构成任何投资建议。交易有风险,入市需谨慎。"
## 分析准则(务必遵守)
0. **只输出结论,不输出思考过程**:禁止复述你的分析步骤。不要写"我先看...""基于上述数据我认为"——直接给结论。
1. **数据说话**:每个判断引用具体排名/涨幅数值,严禁空泛套话("强势"必须改成"连续 4 日稳居前 5,均涨 +4.2%")。
2. **诚实中立**:数据不支持的结论就直言"信号不足,暂无法判断",不要硬凑。
3. **区分资金性质**:这是本分析的核心价值——机构 vs 游资的判断必须基于排名稳定性(标准差),不要凭感觉。
4. **不重复数字**:正文负责解读信号含义,不要照抄罗列已提供的全部原始数据。
5. **简明实战**:总字数 1000-1800 字,重在可执行。
6. **客观推断**:若无明确消息,从量价异动推断可能逻辑并给结论,不要标注"[推断]"或编造具体新闻。
现在请基于下方概念轮动数据进行分析。"""
# ================================================================
# 预计算: 把排名矩阵转成结构化轮动信号
# ================================================================
# 每类信号最多取多少个概念喂给 AI (控制 token)
_TOP_N = 8
def _compute_rotation_signals(dates: list[str], columns: dict) -> dict:
"""从概念涨幅排名矩阵计算轮动信号。
Args:
dates: 日期列表 (最新在最前, 与 columns key 一致)
columns: {日期: [[概念, 涨幅], ...]} 每列各自降序
Returns:
{
"persistent_leaders": [...], # 连续多日稳居前列 (主线)
"rising": [...], # 排名快速跃升 (新晋)
"fading": [...], # 从高位滑落 (退潮)
"institutional": [...], # 排名稳定 (机构特征)
"hot_money": [...], # 排名波动大 (游资特征)
}
每项含: concept, ranks (按 dates 顺序), pcts, avg_rank, rank_std
ranks 时间方向: ranks[0] = 最早日, ranks[-1] = 最新日 (已反转, 左老右新)
"""
if not dates or not columns:
return {}
# 按时间正序 (左老右新) 处理
dates_asc = list(reversed(dates))
# 收集每个概念在各日期的 (排名, 涨幅)。排名 = 该日在列中的索引 + 1。
concept_data: dict[str, list[tuple[int, float]]] = {}
for d in dates_asc:
col = columns.get(d) or []
for idx, (name, pct) in enumerate(col):
concept_data.setdefault(name, []).append((idx + 1, pct))
n_dates = len(dates_asc)
def _stats(ranks_pcts: list[tuple[int, float]]) -> dict:
ranks = [r for r, _ in ranks_pcts]
pcts = [p for _, p in ranks_pcts]
avg = sum(ranks) / len(ranks) if ranks else 0
var = sum((r - avg) ** 2 for r in ranks) / len(ranks) if ranks else 0
return {
"ranks": ranks,
"pcts": [round(p, 4) for p in pcts],
"avg_rank": round(avg, 1),
"rank_std": round(math.sqrt(var), 1),
}
persistent: list[dict] = []
rising: list[dict] = []
fading: list[dict] = []
institutional: list[dict] = []
hot_money: list[dict] = []
for concept, rp in concept_data.items():
# 缺失日补 (大排名, 0 涨幅) 保持时间轴对齐
if len(rp) < n_dates:
rp = rp + [(999, 0.0)] * (n_dates - len(rp))
s = _stats(rp)
s["concept"] = concept
ranks = s["ranks"]
latest_rank = ranks[-1]
earliest_rank = ranks[0]
# 最近 3 日 (不足则全部) 均排名, 判断近期强度
recent = ranks[-min(3, len(ranks)):]
recent_avg = sum(recent) / len(recent)
# 主线: 近期稳居前 10
if recent_avg <= 10 and latest_rank <= 10:
persistent.append(s)
# 新晋: 早期排名靠后(>30), 最新冲进前 20, 跃升幅度大
jump = earliest_rank - latest_rank
if earliest_rank > 30 and latest_rank <= 20 and jump >= 20:
rising.append(s)
# 退潮: 早期排名靠前(<=10), 最新滑落到 30 外
drop = latest_rank - earliest_rank
if earliest_rank <= 10 and latest_rank > 30 and drop >= 20:
fading.append(s)
# 机构: 排名标准差小且平均排名靠前 (稳定强势)
if s["rank_std"] <= 5 and s["avg_rank"] <= 20:
institutional.append(s)
# 游资: 排名标准差大 (波动剧烈)
if s["rank_std"] >= 20:
hot_money.append(s)
# 排序: 主线按近期排名升序; 新晋按跃升幅度降序; 退潮按跌幅降序
persistent.sort(key=lambda x: x["avg_rank"])
rising.sort(key=lambda x: x["ranks"][0] - x["ranks"][-1], reverse=True)
fading.sort(key=lambda x: x["ranks"][-1] - x["ranks"][0], reverse=True)
institutional.sort(key=lambda x: (x["rank_std"], x["avg_rank"]))
hot_money.sort(key=lambda x: x["rank_std"], reverse=True)
return {
"persistent_leaders": persistent[:_TOP_N],
"rising": rising[:_TOP_N],
"fading": fading[:_TOP_N],
"institutional": institutional[:_TOP_N],
"hot_money": hot_money[:_TOP_N],
}
# ================================================================
# Prompt 构建
# ================================================================
def _fmt_pct(v) -> str:
if v is None:
return ""
return f"{v*100:+.2f}%"
def _build_market_block(overview: dict) -> str:
"""大盘背景精简块 (复用 market_overview 已算好的字段)。"""
indices = overview.get("indices") or []
emo = overview.get("emotion") or {}
lim = overview.get("limit") or {}
amt = overview.get("amount") or {}
idx_lines = []
for idx in indices[:4]:
name = idx.get("name") or idx.get("symbol") or "?"
chg = idx.get("change_pct")
idx_lines.append(f"{name} {_fmt_pct(chg)}")
idx_str = " / ".join(idx_lines) or "指数缺失"
total_amount = (amt.get("total") or 0) / 1e8 # 元 → 亿
return (
f"- 指数: {idx_str}\n"
f"- 情绪: {emo.get('score', 50)} ({emo.get('label', '')})\n"
f"- 涨停/炸板/跌停: {lim.get('limit_up', 0)} / {lim.get('broken', 0)} / {lim.get('limit_down', 0)}"
f" (最高连板 {lim.get('max_boards', 0)})\n"
f"- 两市成交额: {total_amount:.0f} 亿元"
)
def _build_signal_block(title: str, items: list[dict]) -> str:
"""轮动信号块: 把预计算的概念信号转成紧凑文本。"""
if not items:
return f"### {title}\n(本类无明显信号)"
lines = [f"### {title}"]
for it in items:
ranks_str = "".join(str(r) if r < 999 else "" for r in it["ranks"])
avg_pct = sum(it["pcts"]) / len(it["pcts"]) if it["pcts"] else 0
lines.append(
f"- {it['concept']}: 排名 {ranks_str} | 均排名 {it['avg_rank']} "
f"| 排名波动σ {it['rank_std']} | 区间均涨 {_fmt_pct(avg_pct)}"
)
return "\n".join(lines)
def _build_user_prompt(signals: dict, overview: dict, days: int, dates: list[str], focus: str) -> str:
"""组装 user 消息: 大盘背景 + 轮动信号 + focus。"""
dates_asc = list(reversed(dates))
date_range = f"{dates_asc[0]} ~ {dates_asc[-1]}" if dates_asc else ""
parts = [
f"# 概念涨幅轮动数据 (最近 {days} 个交易日: {date_range})",
"",
"## 大盘背景",
_build_market_block(overview),
"",
"## 轮动信号 (排名时间方向: 左→右 = 旧→新, 排名越小越强)",
"",
_build_signal_block("🎯 主线 (连续霸榜)", signals.get("persistent_leaders", [])),
"",
_build_signal_block("🆕 新晋强势 (排名跃升)", signals.get("rising", [])),
"",
_build_signal_block("📉 退潮预警 (高位滑落)", signals.get("fading", [])),
"",
_build_signal_block("🏛️ 机构特征 (排名稳定)", signals.get("institutional", [])),
"",
_build_signal_block("🎰 游资特征 (排名波动大)", signals.get("hot_money", [])),
]
if focus.strip():
parts.extend(["", f"## 用户关注点\n{focus.strip()}"])
return "\n".join(parts)
def _build_summary(signals: dict) -> str:
"""meta 事件的摘要 (前端可立即展示)。"""
leaders = signals.get("persistent_leaders", [])
rising = signals.get("rising", [])
fading = signals.get("fading", [])
leader_names = "".join(it["concept"] for it in leaders[:3]) or "暂无明确主线"
return f"主线: {leader_names} | 新晋 {len(rising)} | 退潮 {len(fading)}"
# ================================================================
# 流式主入口
# ================================================================
async def analyze_rotation_stream(
repo,
days: int = 12,
focus: str = "",
quote_service=None,
depth_service=None,
) -> AsyncIterator[str]:
"""流式概念轮动分析: yield 出每个 NDJSON 事件。
Args:
repo: KlineRepository (必填)。
days: 分析最近 N 个交易日 (7-30)。
focus: 用户追加的关注点。
quote_service / depth_service: 可选, 大盘背景装配依赖。
"""
from app.services.rps_rotation import build_rps_rotation
from app.services.market_overview_builder import build_market_overview
# 1. 取轮动矩阵
rotation = build_rps_rotation(repo, days)
dates = rotation.get("dates") or []
columns = rotation.get("columns") or {}
if not dates or not columns:
yield json.dumps({
"type": "error",
"message": "暂无概念轮动数据,请先在「概念分析」页获取概念数据源",
}, ensure_ascii=False)
return
# 2. 预计算轮动信号
signals = _compute_rotation_signals(dates, columns)
# 3. 大盘背景 (失败不阻断, 降级为空)
try:
overview = build_market_overview(repo, quote_service, depth_service)
except Exception as e: # noqa: BLE001
logger.warning("rotation analyze: 大盘背景获取失败, 降级为空: %s", e)
overview = {}
# 4. meta 事件
yield json.dumps({
"type": "meta",
"days": days,
"summary": _build_summary(signals),
}, ensure_ascii=False)
# 5. 构建 prompt + 流式调用 LLM
try:
from app.services.ai_provider import stream_ai_text, ai_configured
if not ai_configured():
yield json.dumps({
"type": "error",
"message": "AI 未配置,请在「设置」页填写 API Key 与接口地址",
}, ensure_ascii=False)
return
user_prompt = _build_user_prompt(signals, overview, days, dates, focus)
async for delta in stream_ai_text(
[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
],
temperature=0.5,
max_tokens=4000,
):
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
except Exception as e: # noqa: BLE001
logger.exception("AI concept rotation analyze failed: %s", e)
yield json.dumps({"type": "error", "message": f"AI 轮动分析失败: {e}"}, ensure_ascii=False)
yield json.dumps({"type": "done"}, ensure_ascii=False)
+586
View File
@@ -0,0 +1,586 @@
"""五档盘口 sealed(真假涨停/跌停) 服务 — 独立旁路线。
架构(完全解耦):
- 只读 enriched(拿涨跌停名单), 不写回 enriched(14列不动)
- sealed 存独立 parquet(data/depth5/date=xxx/part.parquet)
- limit_ladder API 查询时 LEFT JOIN(同 ext_columns 机制)
- signal_limit_up 永远是"价格涨停", sealed 是叠加的真假判定层
数据流:
盘中轮询线程(交易时段, 独立 sleep, 不绑行情轮询):
读 enriched 内存缓存(线程安全) → 涨跌停名单 → tf.depth.batch
→ 算 sealed → 更新内存缓存(不落盘) → sealed_ready=True
盘后定版 job(可配置时间, 默认15:02):
最后拉一次 → 落盘 depth5 parquet(定版)
三层防护节流("设过大设上限, 设过小设最小值"):
① 套餐范围 clamp: Pro 10~120s, Expert 3~300s
② 限速安全 clamp: safe = 60/((rpm*0.8)/batches), 涨跌停多就自动放慢
③ 系统接管通知: 用户设置会超限时, 推 toast 告知已自动调整
"""
from __future__ import annotations
import logging
import math
import threading
import time
from datetime import date, datetime, time as dt_time
from pathlib import Path
import polars as pl
logger = logging.getLogger(__name__)
# 套餐 → (轮询间隔下限s, 上限s)
TIER_INTERVAL_RANGE: dict[str, tuple[float, float]] = {
"pro": (10.0, 120.0),
"expert": (3.0, 300.0),
}
# 兜底: 其他有 DEPTH5_BATCH 的套餐按 pro 范围
DEFAULT_RANGE = (10.0, 120.0)
# 限速余量: 只用 rpm 的 80%, 给系统其他 depth 调用留空间
RPM_MARGIN = 0.8
# 间隔硬下限/上限(任何套餐)
INTERVAL_HARD_MIN = 10.0
INTERVAL_HARD_MAX = 300.0
class DepthService:
"""五档盘口 sealed 服务 — 单例。"""
def __init__(self) -> None:
self._lock = threading.Lock()
self._running = False
self._thread: threading.Thread | None = None
self._repo = None # 延迟注入(KlineRepository)
self._app_state = None # 延迟注入(FastAPI app.state)
# 内存缓存: {symbol: SealedEntry}
# SealedEntry = {sealed_up, sealed_down, ask1_vol, bid1_vol, status, fetched_ts}
self._sealed_cache: dict[str, dict] = {}
self._sealed_ready = False
self._sealed_date: date | None = None # sealed 数据对应的交易日(可能是昨天,如休市)
self._sealed_fetched_ts: float = 0.0 # 上次拉取的 perf_counter
self._sealed_fetched_at: float = 0.0 # 上次拉取的 wall-clock 时间戳
self._persisted_date: date | None = None # 已落盘的日期
# 系统接管状态(防通知刷屏)
self._last_taken_over: bool | None = None
self._last_user_interval: float | None = None
# ================================================================
# 注入
# ================================================================
def set_repo(self, repo) -> None:
self._repo = repo
def set_app_state(self, app_state) -> None:
self._app_state = app_state
# ================================================================
# 生命周期
# ================================================================
def boot_check(self) -> None:
"""启动补跑: 当天 depth5 文件不存在则 finalize 一次; 已存在则恢复内存缓存。"""
if not self._has_capability():
logger.info("depth sealed: 无 DEPTH5_BATCH 能力, 跳过启动补跑")
return
today = date.today()
if self._persisted_for_date(today):
# parquet 已存在: 恢复内存缓存(避免重启后每次查询都读 parquet)
self._restore_from_parquet(today)
return
logger.info("depth sealed: 启动补跑今天定版")
try:
self.finalize()
except Exception as e: # noqa: BLE001
logger.warning("depth sealed 启动补跑失败: %s", e)
def _restore_from_parquet(self, d: date) -> None:
"""从 parquet 恢复内存缓存(服务重启后)。"""
if not self._repo:
return
out = self._repo.store.data_dir / "depth5" / f"date={d.isoformat()}" / "part.parquet"
if not out.exists():
return
try:
df = pl.read_parquet(out)
cache: dict[str, dict] = {}
for row in df.to_dicts():
sym = row.get("symbol")
if not sym:
continue
cache[sym] = {
"sealed_up": row.get("sealed_up"),
"sealed_down": row.get("sealed_down"),
"ask1_vol": row.get("ask1_vol"),
"bid1_vol": row.get("bid1_vol"),
"status": row.get("status"),
"fetched_ts": row.get("fetched_at"),
}
with self._lock:
self._sealed_cache = cache
self._sealed_ready = True
self._sealed_date = d
self._persisted_date = d
logger.info("depth sealed: 从 parquet 恢复 %d 只 (日期=%s)", len(cache), d)
except Exception as e: # noqa: BLE001
logger.warning("depth sealed 从 parquet 恢复失败: %s", e)
def start_polling(self) -> None:
"""启动盘中轮询线程(连板梯队监控开启 + 有能力 + 交易时段)。"""
if self._running:
return
if not self._has_capability():
return
from app.services import preferences
if not preferences.get_limit_ladder_monitor_enabled():
return
self._running = True
self._thread = threading.Thread(target=self._poll_loop, daemon=True)
self._thread.start()
logger.info("depth sealed 盘中轮询已启动")
def stop_polling(self) -> None:
"""停止盘中轮询线程。"""
self._running = False
if self._thread:
self._thread.join(timeout=10)
self._thread = None
logger.info("depth sealed 盘中轮询已停止")
def apply_monitor_toggle(self, enabled: bool) -> None:
"""连板梯队监控开关切换时调用: 开启→启动轮询, 关闭→停止轮询。"""
if enabled:
self.start_polling()
else:
self.stop_polling()
def run_once(self) -> dict:
"""手动触发一次修正(立即拉取 depth + 更新内存缓存)。
不受监控开关限制 — 用户可随时手动修正一次。
返回 {"ok": bool, "count": int, "msg": str}
"""
if not self._has_capability():
return {"ok": False, "count": 0, "msg": "无五档盘口能力(需 Pro+)"}
try:
self._fetch_and_seal(persist=True) # 落盘, 刷新页面不丢
with self._lock:
count = len(self._sealed_cache)
return {"ok": True, "count": count, "msg": f"已修正 {count}"}
except Exception as e: # noqa: BLE001
logger.warning("depth run_once 失败: %s", e)
return {"ok": False, "count": 0, "msg": f"修正失败: {e}"}
# ================================================================
# 核心拉取
# ================================================================
def _fetch_and_seal(self, persist: bool = False) -> None:
"""拉一次 depth.batch, 算 sealed, 更新内存缓存(可选落盘)。
persist=True: 盘后定版, 写 depth5 parquet
persist=False: 盘中轮询, 只更新内存缓存
"""
if not self._repo:
return
# 只读 enriched 内存缓存(线程安全, 避免和 quote_service 写盘竞态)
enriched, enriched_date = self._repo.get_enriched_latest()
if enriched.is_empty():
return
# 筛涨跌停名单(用 fill_null 防止列缺失)
syms_up: list[str] = []
syms_down: list[str] = []
if "signal_limit_up" in enriched.columns:
syms_up = enriched.filter(
pl.col("signal_limit_up").fill_null(False)
)["symbol"].to_list()
if "signal_limit_down" in enriched.columns:
syms_down = enriched.filter(
pl.col("signal_limit_down").fill_null(False)
)["symbol"].to_list()
all_syms = list(dict.fromkeys(syms_up + syms_down)) # 去重保序
if not all_syms:
logger.debug("depth sealed: 当日无涨跌停股, 跳过")
return
# 拉 depth(涨跌停一次拉, 按 capset batch 切片)
depth_data = self._call_depth_batch(all_syms)
if not depth_data:
logger.warning("depth sealed: depth.batch 返回空")
return
up_set = set(syms_up)
down_set = set(syms_down)
now_perf = time.perf_counter()
now_wall = time.time()
new_cache: dict[str, dict] = {}
for sym, d in depth_data.items():
ask_vols = d.get("ask_volumes") or []
bid_vols = d.get("bid_volumes") or []
ask1 = ask_vols[0] if ask_vols else None
bid1 = bid_vols[0] if bid_vols else None
# depth 返回的 timestamp(毫秒 epoch), 回退到当前 wall-clock
depth_ts = d.get("timestamp")
fetched = (depth_ts / 1000.0) if isinstance(depth_ts, (int, float)) and depth_ts else now_wall
entry = {
# 涨停真封: 涨停价上卖一(主动卖压)为 0
"sealed_up": (ask1 == 0) if sym in up_set and ask1 is not None else None,
# 跌停真封: 跌停价上买一为 0
"sealed_down": (bid1 == 0) if sym in down_set and bid1 is not None else None,
"ask1_vol": ask1,
"bid1_vol": bid1,
"status": "limit_down" if sym in down_set and sym not in up_set else "limit_up",
"fetched_ts": fetched,
}
new_cache[sym] = entry
with self._lock:
self._sealed_cache = new_cache
self._sealed_ready = True
self._sealed_date = enriched_date # 记录数据对应的交易日(可能是昨天,如休市)
self._sealed_fetched_ts = now_perf
self._sealed_fetched_at = now_wall
logger.info("depth sealed: 拉取 %d 只 (涨停%d/跌停%d) 日期=%s%s",
len(new_cache), len(syms_up), len(syms_down),
enriched_date, " → 落盘" if persist else "")
# 缓存已更新: 通知 SSE 推 depth_updated, 触发连板梯队刷新封单数据。
self._notify_depth_updated(len(new_cache))
if persist and enriched_date:
self._persist(enriched_date)
def _call_depth_batch(self, symbols: list[str]) -> dict:
"""调 tf.depth.batch, 按 capset 的 batch 切片 + 节流。返回 {symbol: MarketDepth}。"""
from app.tickflow.client import get_client
tf = get_client()
capset = self._get_capset()
lim = capset.limits(__import__("app.tickflow.capabilities", fromlist=["Cap"]).Cap.DEPTH5_BATCH)
batch_size = (lim.batch if lim and lim.batch else 100)
rpm = (lim.rpm if lim and lim.rpm else 30)
# 批间隔 = 60/rpm(匀速)
inter_batch = 60.0 / rpm if rpm > 0 else 2.0
result: dict = {}
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
for i, chunk in enumerate(chunks):
if i > 0:
time.sleep(inter_batch)
try:
# SDK 的 batch 内部已按 batch_size 切, 这里再切一层防单请求过大
data = tf.depth.batch(chunk)
if isinstance(data, dict):
result.update(data)
except Exception as e: # noqa: BLE001
logger.warning("depth.batch 第 %d 批失败(%d 只): %s", i + 1, len(chunk), e)
# 单批失败不影响其他批
return result
def finalize(self) -> None:
"""盘后定版: 拉一次 + 落盘。"""
if not self._has_capability():
return
self._fetch_and_seal(persist=True)
# ================================================================
# 落盘
# ================================================================
def _persist(self, today: date) -> None:
"""把内存缓存写 depth5/date=今天/part.parquet。"""
with self._lock:
cache = dict(self._sealed_cache)
if not cache:
return
rows = []
for sym, e in cache.items():
rows.append({
"symbol": sym,
"sealed_up": e.get("sealed_up"),
"sealed_down": e.get("sealed_down"),
"ask1_vol": e.get("ask1_vol"),
"bid1_vol": e.get("bid1_vol"),
"status": e.get("status"),
"fetched_at": e.get("fetched_ts"),
})
# 显式 schema: sealed_up/sealed_down 是 bool 与 None 混合, 不指定 schema
# polars 会按首行推断类型, 后续遇到不一致 (bool vs null) 报
# "could not append value: false of type: bool to the builder"。
df = pl.DataFrame(rows, schema={
"symbol": pl.Utf8,
"sealed_up": pl.Boolean,
"sealed_down": pl.Boolean,
"ask1_vol": pl.Int64,
"bid1_vol": pl.Int64,
"status": pl.Utf8,
"fetched_at": pl.Float64,
})
ds = today.isoformat()
out = self._repo.store.data_dir / "depth5" / f"date={ds}" / "part.parquet"
out.parent.mkdir(parents=True, exist_ok=True)
df.write_parquet(out)
self._persisted_date = today
logger.info("depth sealed 落盘: %d 行 → %s", df.height, out)
def _persisted_for_date(self, d: date) -> bool:
"""检查某日 depth5 文件是否已存在。"""
if not self._repo:
return False
out = self._repo.store.data_dir / "depth5" / f"date={d.isoformat()}" / "part.parquet"
return out.exists()
# ================================================================
# 查询(供 limit_ladder API 用)
# ================================================================
def get_sealed_map(self, target_date: date, is_down: bool) -> dict:
"""返回 {symbol: {sealed, vol, ready, age}} 供 JOIN。
优先内存缓存(盘中), 回退 parquet(历史/盘后)。
sealed: bool | None (None=待确认或降级)
vol: 封单量(int) | None
ready: sealed 数据是否就绪(False→降级标识)
age: 距上次拉取秒数(盘后定版为 None)
"""
# 内存缓存(sealed 数据对应的交易日 = target_date 时才用)
if self._sealed_date and target_date == self._sealed_date and self._sealed_ready and self._sealed_cache:
return self._read_from_memory(is_down)
# parquet(历史或盘后定版)
return self._read_from_parquet(target_date, is_down)
def _read_from_memory(self, is_down: bool) -> dict:
sealed_key = "sealed_down" if is_down else "sealed_up"
# 封单量: 涨停=买一量(涨停价买单堆积), 跌停=卖一量(跌停价卖单堆积)
vol_key = "ask1_vol" if is_down else "bid1_vol"
now = time.perf_counter()
with self._lock:
cache = dict(self._sealed_cache)
fetched_ts = self._sealed_fetched_ts
age = (now - fetched_ts) if fetched_ts else 0.0
result = {}
for sym, e in cache.items():
result[sym] = {
"sealed": e.get(sealed_key),
"vol": e.get(vol_key),
"ready": True,
"age": age,
}
return result
def _read_from_parquet(self, target_date: date, is_down: bool) -> dict:
if not self._repo:
return {}
out = self._repo.store.data_dir / "depth5" / f"date={target_date.isoformat()}" / "part.parquet"
if not out.exists():
return {}
try:
df = pl.read_parquet(out)
except Exception as e: # noqa: BLE001
logger.warning("depth5 parquet 读取失败: %s", e)
return {}
sealed_key = "sealed_down" if is_down else "sealed_up"
# 封单量: 涨停=买一量, 跌停=卖一量
vol_key = "ask1_vol" if is_down else "bid1_vol"
result = {}
for row in df.to_dicts():
sym = row.get("symbol")
if not sym:
continue
result[sym] = {
"sealed": row.get(sealed_key),
"vol": row.get(vol_key),
"ready": True,
"age": None, # 盘后定版, 无 age
}
return result
def is_sealed_ready(self, target_date: date) -> bool:
"""sealed 数据是否就绪(供前端降级判定)。"""
# 内存缓存对应的数据日 == 查询日 → 看内存就绪状态
if self._sealed_date and target_date == self._sealed_date:
return self._sealed_ready
# 其他日期: 有 parquet 就 ready
return self._persisted_for_date(target_date)
def get_sealed_age(self, target_date: date) -> float | None:
"""返回 sealed 数据 age(秒), 盘后定版为 None。"""
if self._sealed_date and target_date == self._sealed_date and self._sealed_ready and self._sealed_fetched_ts:
return time.perf_counter() - self._sealed_fetched_ts
return None
# ================================================================
# 盘中轮询线程
# ================================================================
def _poll_loop(self) -> None:
"""盘中轮询: 按 capset 自适应间隔拉 depth, 更新内存缓存。"""
while self._running:
try:
if self._is_trading_hours():
self._poll_once()
else:
logger.debug("depth sealed: 非交易时段, 跳过")
except Exception as e: # noqa: BLE001
logger.warning("depth sealed 轮询异常: %s", e)
# 等待下一轮(用 _running 检查保证能及时退出)
interval = self._current_sleep_interval()
waited = 0.0
while self._running and waited < interval:
time.sleep(0.5)
waited += 0.5
def _poll_once(self) -> None:
"""单次轮询: 算间隔(三层防护) → 拉取 → 检测系统接管通知。"""
# 数当前涨跌停股
n = self._count_limit_stocks()
if n == 0:
return
interval, taken_over, user_interval = self._compute_interval(n)
# 系统接管通知(状态切换时才推, 防刷屏)
if taken_over and (self._last_taken_over is False or self._last_user_interval != user_interval):
self._notify_takeover(n, user_interval, interval)
self._last_taken_over = taken_over
self._last_user_interval = user_interval
self._fetch_and_seal(persist=False)
def _current_sleep_interval(self) -> float:
"""计算当前 sleep 间隔(供 _poll_loop 等待用)。"""
n = self._count_limit_stocks()
if n == 0:
return 30.0 # 无涨跌停, 慢轮询
interval, _, _ = self._compute_interval(n)
return interval
# ================================================================
# 三层防护节流
# ================================================================
def _compute_interval(self, n_symbols: int) -> tuple[float, bool, float]:
"""三层防护计算实际轮询间隔。
返回 (actual_interval, taken_over, user_interval)
- actual_interval: 实际使用的间隔(秒)
- taken_over: 是否被系统接管(用户设置会超限)
- user_interval: 用户设置(经套餐 clamp 后)的间隔
"""
from app.services import preferences
from app.tickflow.policy import tier_label
capset = self._get_capset()
lim = capset.limits(__import__("app.tickflow.capabilities", fromlist=["Cap"]).Cap.DEPTH5_BATCH)
batch_size = (lim.batch if lim and lim.batch else 100)
rpm = (lim.rpm if lim and lim.rpm else 30)
# ① 套餐范围 clamp
tier = tier_label().split()[0].split("+")[0].strip().lower()
lo, hi = TIER_INTERVAL_RANGE.get(tier, DEFAULT_RANGE)
raw_user = preferences.get_depth_polling_interval()
user_interval = max(lo, min(hi, raw_user))
# ② 限速安全 clamp
batches = max(1, math.ceil(n_symbols / batch_size))
usable_rpm = rpm * RPM_MARGIN
calls_per_min = usable_rpm / batches if batches > 0 else usable_rpm
safe_interval = 60.0 / calls_per_min if calls_per_min > 0 else INTERVAL_HARD_MAX
# 实际: 取用户设置和安全的较大值
actual = max(user_interval, safe_interval)
# 硬上下限
actual = max(INTERVAL_HARD_MIN, min(actual, INTERVAL_HARD_MAX))
taken_over = safe_interval > user_interval
return actual, taken_over, user_interval
def _count_limit_stocks(self) -> int:
"""数当前涨跌停股总数(供节流计算)。"""
if not self._repo:
return 0
enriched, _ = self._repo.get_enriched_latest()
if enriched.is_empty():
return 0
n = 0
if "signal_limit_up" in enriched.columns:
n += enriched.filter(pl.col("signal_limit_up").fill_null(False)).height
if "signal_limit_down" in enriched.columns:
n += enriched.filter(pl.col("signal_limit_down").fill_null(False)).height
return n
# ================================================================
# 通知
# ================================================================
def _notify_takeover(self, n_stocks: int, user_interval: float, actual_interval: float) -> None:
"""系统接管通知: 复用 quote_service 的 _pending_alerts 通道。"""
if not self._app_state:
return
qs = getattr(self._app_state, "quote_service", None)
if not qs:
return
msg = (f"五档轮询: 当前涨跌停 {n_stocks} 只, 您设置的 {user_interval:.0f} 秒间隔会超限, "
f"系统已自动调整为 {actual_interval:.0f}")
alert = {
"source": "depth",
"type": "takeover",
"message": msg,
}
try:
with qs._lock:
qs._pending_alerts.append(alert)
qs._alert_event.set()
except Exception as e: # noqa: BLE001
logger.debug("depth 接管通知推送失败: %s", e)
def _notify_depth_updated(self, count: int) -> None:
"""修正完成通知: set quote_service._depth_update_event, SSE 推 depth_updated 刷新连板梯队。"""
if not self._app_state:
return
qs = getattr(self._app_state, "quote_service", None)
if not qs:
return
try:
qs.notify_depth_updated()
except Exception as e: # noqa: BLE001
logger.debug("depth 更新通知推送失败: %s", e)
# ================================================================
# 工具
# ================================================================
def _has_capability(self) -> bool:
capset = self._get_capset()
from app.tickflow.capabilities import Cap
return capset.has(Cap.DEPTH5_BATCH)
def _get_capset(self):
"""获取当前 capset(优先 app.state, 回退 detect)。"""
if self._app_state:
cs = getattr(self._app_state, "capabilities", None)
if cs:
return cs
from app.tickflow.policy import detect_capabilities
return detect_capabilities()
@staticmethod
def _is_trading_hours() -> bool:
now = datetime.now()
t = now.time()
morning = dt_time(9, 25) <= t <= dt_time(11, 35)
afternoon = dt_time(12, 55) <= t <= dt_time(15, 5)
return now.weekday() < 5 and (morning or afternoon)
+521
View File
@@ -0,0 +1,521 @@
"""扩展数据服务 — 配置管理 + 文件解析 + Parquet 存储。"""
from __future__ import annotations
import json
import logging
from datetime import date, datetime
from pathlib import Path
from typing import Literal
import polars as pl
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# 配置模型
# ---------------------------------------------------------------------------
class ExtField:
"""扩展字段定义。"""
__slots__ = ("name", "dtype", "label")
def __init__(self, name: str, dtype: str = "string", label: str = "") -> None:
self.name = name
self.dtype = dtype # string | int | float | bool
self.label = label or name
def to_dict(self) -> dict:
return {"name": self.name, "dtype": self.dtype, "label": self.label}
@classmethod
def from_dict(cls, d: dict) -> ExtField:
return cls(d["name"], d.get("dtype", "string"), d.get("label", ""))
class PullConfig:
"""定时拉取配置。"""
__slots__ = (
"url", "method", "headers", "body", "response_path",
"field_map", "schedule_minutes", "enabled",
"last_run", "last_status", "last_message", "last_rows",
"next_run",
)
def __init__(
self,
url: str = "",
method: str = "GET",
headers: dict[str, str] | None = None,
body: str | None = None,
response_path: str = "",
field_map: dict[str, str] | None = None,
schedule_minutes: int = 1440,
enabled: bool = False,
last_run: str | None = None,
last_status: str | None = None,
last_message: str | None = None,
last_rows: int | None = None,
next_run: str | None = None,
) -> None:
self.url = url
self.method = method # GET | POST
self.headers = headers or {}
self.body = body # JSON string (POST body template)
self.response_path = response_path # dot-path to rows array, e.g. "data.list"
self.field_map = field_map or {} # external_name → config_field_name
self.schedule_minutes = schedule_minutes
self.enabled = enabled
self.last_run = last_run
self.last_status = last_status # "success" | "error"
self.last_message = last_message
self.last_rows = last_rows
self.next_run = next_run # 下次预计运行 (ISO, 调度器写入)
def to_dict(self) -> dict:
return {
"url": self.url,
"method": self.method,
"headers": self.headers,
"body": self.body,
"response_path": self.response_path,
"field_map": self.field_map,
"schedule_minutes": self.schedule_minutes,
"enabled": self.enabled,
"last_run": self.last_run,
"last_status": self.last_status,
"last_message": self.last_message,
"last_rows": self.last_rows,
"next_run": self.next_run,
}
@classmethod
def from_dict(cls, d: dict) -> PullConfig:
if not d:
return cls()
return cls(
url=d.get("url", ""),
method=d.get("method", "GET"),
headers=d.get("headers"),
body=d.get("body"),
response_path=d.get("response_path", ""),
field_map=d.get("field_map"),
schedule_minutes=d.get("schedule_minutes", 1440),
enabled=d.get("enabled", False),
last_run=d.get("last_run"),
last_status=d.get("last_status"),
last_message=d.get("last_message"),
last_rows=d.get("last_rows"),
next_run=d.get("next_run"),
)
class ExtConfig:
"""一个扩展数据源的完整配置。"""
__slots__ = (
"id", "label", "mode", "fields", "description",
"symbol_map", "code_map",
"created_at", "updated_at", "pull",
)
def __init__(
self,
id: str,
label: str,
mode: Literal["snapshot", "timeseries"],
fields: list[ExtField],
description: str = "",
symbol_map: dict | None = None,
code_map: dict | None = None,
created_at: str | None = None,
updated_at: str | None = None,
pull: PullConfig | None = None,
) -> None:
self.id = id
self.label = label
self.mode = mode
self.fields = fields
self.description = description
# 映射关系: {"type": "mapped", "col": "原始列名"} 或 {"type": "computed", "from": "symbol|code", "method": "strip_exchange|append_exchange"}
self.symbol_map = symbol_map or {}
self.code_map = code_map or {}
self.created_at = created_at or datetime.now().isoformat()
self.updated_at = updated_at or datetime.now().isoformat()
self.pull = pull
def to_dict(self) -> dict:
d = {
"id": self.id,
"label": self.label,
"mode": self.mode,
"fields": [f.to_dict() for f in self.fields],
"description": self.description,
"symbol_map": self.symbol_map,
"code_map": self.code_map,
"created_at": self.created_at,
"updated_at": self.updated_at,
}
if self.pull:
d["pull"] = self.pull.to_dict()
return d
@classmethod
def from_dict(cls, d: dict) -> ExtConfig:
return cls(
id=d["id"],
label=d["label"],
mode=d["mode"],
fields=[ExtField.from_dict(f) for f in d.get("fields", [])],
description=d.get("description", ""),
symbol_map=d.get("symbol_map"),
code_map=d.get("code_map"),
created_at=d.get("created_at"),
updated_at=d.get("updated_at"),
pull=PullConfig.from_dict(d["pull"]) if d.get("pull") else None,
)
# ---------------------------------------------------------------------------
# 配置持久化
# ---------------------------------------------------------------------------
class ExtConfigStore:
"""扩展数据配置文件读写 — 每个表独立目录 data/ext/{config_id}/config.json。"""
def __init__(self, data_dir: Path) -> None:
self._base = data_dir / "ext_data"
def _config_path(self, config_id: str) -> Path:
return self._base / config_id / "config.json"
def load_all(self) -> list[ExtConfig]:
# 兼容旧版: 如果目录为空且旧配置文件存在则迁移
if not self._base.exists() or not any(self._base.iterdir()):
old = self._base.parent / "ext_configs.json"
if not old.exists():
old = self._base.parent / "ext_configs.json.bak"
if old.exists():
self._migrate_legacy(old)
if not self._base.exists():
return []
configs = []
for d in sorted(self._base.iterdir()):
cp = d / "config.json"
if d.is_dir() and cp.exists():
try:
raw = json.loads(cp.read_text(encoding="utf-8"))
configs.append(ExtConfig.from_dict(raw))
except Exception as e:
logger.warning("扩展表配置解析失败 %s: %s", cp, e)
return configs
def get(self, config_id: str) -> ExtConfig | None:
cp = self._config_path(config_id)
if not cp.exists():
return None
try:
raw = json.loads(cp.read_text(encoding="utf-8"))
return ExtConfig.from_dict(raw)
except Exception:
return None
def upsert(self, config: ExtConfig) -> None:
config.updated_at = datetime.now().isoformat()
cp = self._config_path(config.id)
cp.parent.mkdir(parents=True, exist_ok=True)
cp.write_text(
json.dumps(config.to_dict(), ensure_ascii=False, indent=2),
encoding="utf-8",
)
def delete(self, config_id: str) -> bool:
import shutil
cp = self._config_path(config_id)
if not cp.exists():
return False
shutil.rmtree(cp.parent, ignore_errors=True)
return True
def _migrate_legacy(self, old_path: Path) -> None:
"""一次性迁移旧版 ext_configs.json 到独立目录结构。"""
try:
raw = json.loads(old_path.read_text(encoding="utf-8"))
configs = [ExtConfig.from_dict(d) for d in raw]
for c in configs:
cp = self._config_path(c.id)
cp.parent.mkdir(parents=True, exist_ok=True)
cp.write_text(
json.dumps(c.to_dict(), ensure_ascii=False, indent=2),
encoding="utf-8",
)
# 迁移完成后重命名旧文件作为备份
backup = old_path.with_suffix(".json.bak")
old_path.rename(backup)
logger.info("ext_configs.json 已迁移至 ext/ (备份: %s)", backup.name)
except Exception as e:
logger.warning("ext_configs 迁移失败: %s", e)
# ---------------------------------------------------------------------------
# CSV / Excel 解析 → Parquet 写入
# ---------------------------------------------------------------------------
_POLARS_DTYPE_MAP = {
"string": pl.Utf8,
"int": pl.Int64,
"float": pl.Float64,
"bool": pl.Boolean,
}
def build_code_lookup(data_dir: Path) -> dict[str, str]:
"""从 instruments 维表构建 code → symbol 映射。"""
path = data_dir / "instruments" / "instruments.parquet"
if not path.exists():
return {}
try:
df = pl.read_parquet(path, columns=["code", "symbol"])
return dict(zip(df["code"].to_list(), df["symbol"].to_list()))
except Exception:
return {}
def normalize_symbol(series: pl.Series, lookup: dict[str, str] | None = None) -> pl.Series:
"""将 symbol 列标准化为 代码.交易所 格式。
优先使用 instruments 维表查找 code → symbol,确保 100% 准确。
查不到时按规则兜底:6开头 → .SH,其余 → .SZ。
"""
_lookup = lookup or {}
def _fix_one(val: str) -> str:
if not val:
return val
val = val.strip()
# 已经是标准格式(含 .),直接返回
if "." in val:
return val
# 纯6位数字代码 → 优先查维表
if len(val) == 6 and val.isdigit():
mapped = _lookup.get(val)
if mapped:
return mapped
# 兜底规则
if val.startswith(("6",)):
return f"{val}.SH"
else:
return f"{val}.SZ"
return val
return series.map_elements(_fix_one, return_dtype=pl.Utf8)
def ensure_utf8_csv(file_path: Path) -> Path:
"""确保 CSV 文件以 UTF-8 编码可读,非 UTF-8(如 GBK/GB18030)则转换。
国内行情软件(同花顺/东财/通达信)和 Windows 中文 Excel 导出的 CSV 多为
GBK 系编码,Polars 的 read_csv 默认按 UTF-8 解析会抛 "invalid utf-8 sequence"
这里在交给 Polars 前做一次编码规范化。
返回值:若已是 UTF-8 则返回原路径;否则在同目录写一个 *.utf8 文件并返回它
(调用方用临时目录,随目录一起清理)。
"""
raw = file_path.read_bytes()
# BOM 处理:UTF-8-SIG 等带 BOM 文件直接交给 Polars(它认识 BOM)
try:
raw.decode("utf-8")
return file_path # 已是合法 UTF-8
except UnicodeDecodeError:
pass
# 依次尝试常见中文编码,第一个能完整解码的即为命中
for enc in ("gb18030", "gbk", "gb2312", "big5"):
try:
text = raw.decode(enc)
except UnicodeDecodeError:
continue
out_path = file_path.with_suffix(file_path.suffix + ".utf8")
out_path.write_text(text, encoding="utf-8")
logger.info("CSV 编码转换 %s%s (%s)", file_path.name, out_path.name, enc)
return out_path
# 都无法解码:返回原路径,让 Polars 抛出更精确的原始错误
return file_path
def parse_upload_file(file_path: Path, symbol_col: str = "symbol", data_dir: Path | None = None) -> pl.DataFrame:
"""解析上传的 CSV / Excel 文件为 Polars DataFrame。"""
suffix = file_path.suffix.lower()
if suffix == ".csv":
df = pl.read_csv(ensure_utf8_csv(file_path), infer_schema_length=10000)
elif suffix in (".xlsx", ".xls"):
df = pl.read_excel(file_path)
else:
raise ValueError(f"不支持的文件格式: {suffix}")
if symbol_col not in df.columns:
# 尝试模糊匹配
candidates = [c for c in df.columns if c.lower() in ("symbol", "code", "代码", "标的")]
if candidates:
df = df.rename({candidates[0]: symbol_col})
else:
raise ValueError(f"未找到标的代码列 (symbol),可选列: {df.columns}")
# 确保 symbol 列为字符串并标准化
lookup = build_code_lookup(data_dir) if data_dir else None
df = df.with_columns(normalize_symbol(df[symbol_col].cast(pl.Utf8), lookup))
return df
def cast_df_to_schema(df: pl.DataFrame, fields: list[ExtField]) -> pl.DataFrame:
"""按配置的字段类型转换 DataFrame 列类型。"""
for f in fields:
if f.name in df.columns:
target = _POLARS_DTYPE_MAP.get(f.dtype, pl.Utf8)
df = df.with_columns(pl.col(f.name).cast(target))
return df
def _config_dir(config_id: str, data_dir: Path) -> Path:
"""返回扩展配置的根目录 data/ext_data/{config_id}/。"""
return data_dir / "ext_data" / config_id
def write_ext_parquet(
df: pl.DataFrame,
config: ExtConfig,
data_dir: Path,
snapshot_date: date | None = None,
) -> int:
"""将 DataFrame 写入扩展数据 Parquet。
目录结构:
- snapshot: data/ext_data/{id}/part.parquet(与 config.json 同级,覆盖写)
- timeseries: data/ext_data/{id}/timeseries/date=xxx/part.parquet(按日分区)
Returns:
写入行数。
"""
snap = snapshot_date or date.today()
cfg_dir = _config_dir(config.id, data_dir)
# 标准化 symbol 列: 用维表查找 → 准确匹配交易所
if "symbol" in df.columns:
lookup = build_code_lookup(data_dir)
df = df.with_columns(normalize_symbol(df["symbol"], lookup))
if config.mode == "snapshot":
# 快照: 与 config.json 同级,直接覆盖
cfg_dir.mkdir(parents=True, exist_ok=True)
out_path = cfg_dir / "part.parquet"
# 如果已有文件,合并去重后覆盖
if out_path.exists():
try:
existing = pl.read_parquet(out_path)
key = "symbol" if "symbol" in df.columns else df.columns[0]
df = pl.concat([existing, df]).unique(subset=[key], keep="last")
except Exception as e:
# schema 不一致 (列不同) 时 concat 失败 → 直接用新 df 覆盖。
# 记日志而非静默吞掉, 便于排查"数据结构错乱"类问题。
logger.warning("扩展表 %s 合并去重失败, 将覆盖写入: %s", config.id, e)
else:
# 时序: timeseries/ 下按日期分区
out_dir = cfg_dir / "timeseries" / f"date={snap}"
out_dir.mkdir(parents=True, exist_ok=True)
out_path = out_dir / "part.parquet"
# 如果已有文件,合并去重
if out_path.exists():
try:
existing = pl.read_parquet(out_path)
key = "symbol" if "symbol" in df.columns else df.columns[0]
df = pl.concat([existing, df]).unique(subset=[key], keep="last")
except Exception as e:
logger.warning("扩展表 %s 合并去重失败, 将覆盖写入: %s", config.id, e)
df = cast_df_to_schema(df, config.fields)
df.write_parquet(out_path)
logger.info("扩展表写入: %s%s (%d 行)", config.id, out_path, len(df))
return len(df)
def delete_ext_parquet(config_id: str, data_dir: Path) -> None:
"""删除扩展数据源关联的所有 Parquet 数据(保留 config.json)。
- snapshot: 删除 ext_data/{id}/part.parquet
- timeseries: 删除 ext_data/{id}/timeseries/ 目录
"""
cfg_dir = _config_dir(config_id, data_dir)
# 删除快照文件
snap = cfg_dir / "part.parquet"
if snap.exists():
snap.unlink()
# 删除时序目录
ts_dir = cfg_dir / "timeseries"
if ts_dir.exists():
import shutil
shutil.rmtree(ts_dir, ignore_errors=True)
def fix_symbol_format(config: ExtConfig, data_dir: Path) -> int:
"""扫描该扩展配置的所有 Parquet 文件,将 symbol 列标准化为 代码.交易所 格式。
- snapshot: 扫描 ext_data/{id}/part.parquet
- timeseries: 扫描 ext_data/{id}/timeseries/date=xxx/part.parquet
Returns:
修复的文件数。
"""
cfg_dir = _config_dir(config.id, data_dir)
if not cfg_dir.exists():
return 0
# 收集需要扫描的 parquet 文件列表
parquet_files: list[Path] = []
if config.mode == "snapshot":
p = cfg_dir / "part.parquet"
if p.exists():
parquet_files.append(p)
else:
ts_dir = cfg_dir / "timeseries"
if ts_dir.exists():
for part_dir in sorted(ts_dir.iterdir()):
if not part_dir.is_dir() or not part_dir.name.startswith("date="):
continue
p = part_dir / "part.parquet"
if p.exists():
parquet_files.append(p)
fixed = 0
lookup = build_code_lookup(data_dir)
for parquet_path in parquet_files:
try:
df = pl.read_parquet(parquet_path)
if "symbol" not in df.columns:
continue
old = df["symbol"].to_list()
df = df.with_columns(normalize_symbol(df["symbol"], lookup))
new = df["symbol"].to_list()
if old != new:
df.write_parquet(parquet_path)
fixed += 1
logger.info("代码格式修复: %s/%s (%d 行)", config.id, parquet_path.parent.name, len(df))
except Exception as e:
logger.warning("代码格式修复跳过 %s: %s", parquet_path, e)
return fixed
def rows_to_parquet(
rows: list[dict],
config: ExtConfig,
data_dir: Path,
snapshot_date: date | None = None,
) -> int:
"""将 JSON 行列表转为 DataFrame 写入 Parquet,复用 write_ext_parquet 的存储逻辑。
Returns:
写入行数。
"""
df = pl.DataFrame(rows)
if "symbol" in df.columns:
df = df.with_columns(pl.col("symbol").cast(pl.Utf8))
return write_ext_parquet(df, config, data_dir, snapshot_date=snapshot_date)
+236
View File
@@ -0,0 +1,236 @@
"""内置扩展数据预设 — 概念/行业首次启动自动拉取。
设计原则:
- 扩展数据通用逻辑零改动 (ExtConfig / fetch_and_ingest / API / 前端均不动)
- 仅在本模块做「接口结构 → 本地 schema」的转换
- 「已存在则跳过」: 绝不覆盖用户已有数据, 老用户零影响
- 拉取失败只记 warning, 不阻断启动 (保持「没数据也能跑」)
种子数据来源: https://files.688798.xyz/ths/{concepts,industries}.json
作者更新数据只需改接口上的 JSON, 用户下次拉取自动同步, 无需发版。
接入点: app.main.lifespan → ensure_builtin_presets(store.data_dir)
"""
from __future__ import annotations
import logging
from pathlib import Path
from app.services.ext_data import (
ExtConfig,
ExtConfigStore,
ExtField,
PullConfig,
rows_to_parquet,
)
logger = logging.getLogger(__name__)
# 种子数据源 (作者维护, 改这里即对所有用户生效)
_THS_BASE = "https://files.688798.xyz/ths"
# ---------------------------------------------------------------------------
# 预设定义: 字段结构 + 拉取配方
# ---------------------------------------------------------------------------
def _concept_preset() -> ExtConfig:
"""扩展概念 (ext_gn_ths)。
接口结构: [{symbol, name, concepts: [概念1, 概念2, ...]}]
本地 schema: 股票代码 / 股票简称 / 所属概念(分号拼接) / symbol / code
"""
return ExtConfig(
id="ext_gn_ths",
label="扩展概念",
mode="snapshot",
fields=[
ExtField("symbol", "string", "标的代码"),
ExtField("code", "string", "代码"),
ExtField("股票代码", "string", "股票代码"),
ExtField("股票简称", "string", "股票简称"),
ExtField("所属概念", "string", "所属概念"),
],
description="同花顺概念分类 (首次启动自动拉取, 可在扩展数据页手动更新)",
symbol_map={"type": "mapped", "col": "股票代码"},
code_map={"type": "computed", "from": "symbol", "method": "strip_exchange"},
pull=PullConfig(
url=f"{_THS_BASE}/concepts.json",
method="GET",
schedule_minutes=1440,
enabled=False,
),
)
def _industry_preset() -> ExtConfig:
"""扩展行业 (ext_hy_ths)。
接口结构: [{symbol, name, industries: [一级行业, 二级行业, 三级行业]}]
本地 schema: 股票代码 / 股票简称 / 所属同花顺行业(横杠拼接) / symbol / code
"""
return ExtConfig(
id="ext_hy_ths",
label="扩展行业",
mode="snapshot",
fields=[
ExtField("symbol", "string", "标的代码"),
ExtField("code", "string", "代码"),
ExtField("股票代码", "string", "股票代码"),
ExtField("股票简称", "string", "股票简称"),
ExtField("所属同花顺行业", "string", "所属同花顺行业"),
],
description="同花顺行业分类 (首次启动自动拉取, 可在扩展数据页手动更新)",
symbol_map={"type": "mapped", "col": "股票代码"},
code_map={"type": "computed", "from": "symbol", "method": "strip_exchange"},
pull=PullConfig(
url=f"{_THS_BASE}/industries.json",
method="GET",
schedule_minutes=1440,
enabled=False,
),
)
def _presets() -> list[ExtConfig]:
return [_concept_preset(), _industry_preset()]
# ---------------------------------------------------------------------------
# 接口结构 → 本地 schema 转换 (仅预设使用)
# ---------------------------------------------------------------------------
def _symbol_to_code(symbol: str) -> str:
"""symbol (000001.SZ) → code (000001)。"""
return symbol.split(".", 1)[0] if "." in symbol else symbol
def _flatten_concept_rows(raw_rows: list[dict]) -> list[dict]:
"""概念: concepts 数组 → 分号拼接成「所属概念」字符串。
[{symbol, name, concepts:[...]}] → [{股票代码, 股票简称, 所属概念, symbol, code}]
注: code 由 symbol 派生 (000001.SZ → 000001), 因 rows_to_parquet 不执行 code_map。
"""
out: list[dict] = []
for r in raw_rows:
sym = (r.get("symbol") or "").strip()
if not sym:
continue
concepts = r.get("concepts") or []
out.append({
"股票代码": sym,
"股票简称": r.get("name") or "",
"所属概念": ";".join(str(c) for c in concepts if c),
"symbol": sym,
"code": _symbol_to_code(sym),
})
return out
def _flatten_industry_rows(raw_rows: list[dict]) -> list[dict]:
"""行业: industries 数组 → 横杠拼接成「所属同花顺行业」字符串。
[{symbol, name, industries:[...]}] → [{股票代码, 股票简称, 所属同花顺行业, symbol, code}]
"""
out: list[dict] = []
for r in raw_rows:
sym = (r.get("symbol") or "").strip()
if not sym:
continue
inds = r.get("industries") or []
out.append({
"股票代码": sym,
"股票简称": r.get("name") or "",
"所属同花顺行业": "-".join(str(i) for i in inds if i),
"symbol": sym,
"code": _symbol_to_code(sym),
})
return out
# ---------------------------------------------------------------------------
# 拉取执行 (复用 httpx, 不依赖 fetch_and_ingest 的 PullConfig 路径)
# ---------------------------------------------------------------------------
async def _fetch_json(url: str) -> list[dict]:
"""请求 JSON 接口, 返回行数组。超时 30s, 失败抛异常由调用方兜底。"""
import httpx
async with httpx.AsyncClient(timeout=30) as client:
resp = await client.get(url)
resp.raise_for_status()
data = resp.json()
if not isinstance(data, list):
raise ValueError(f"接口返回不是数组: {type(data)}")
return data
async def _seed_one(config: ExtConfig, flatten, data_dir: Path) -> int:
"""拉取 + 转换 + 写入单个预设。返回写入行数。"""
from datetime import date
raw = await _fetch_json(config.pull.url)
rows = flatten(raw)
if not rows:
raise ValueError(f"接口返回 0 行: {config.pull.url}")
n = rows_to_parquet(rows, config, data_dir, snapshot_date=date.today())
return n
# ---------------------------------------------------------------------------
# 对外入口
# ---------------------------------------------------------------------------
def get_preset(config_id: str) -> ExtConfig | None:
"""按 id 取预设定义 (供 API 层校验 id 合法性)。"""
for c in _presets():
if c.id == config_id:
return c
return None
async def ensure_builtin_presets(data_dir: Path) -> None:
"""启动时: 为缺失的预设创建 config.json (含 pull 配置), 但【不拉取数据】。
设计: 数据获取改为用户在概念/行业页手动点「获取数据」触发, 避免启动时
网络请求阻塞, 也避免「自动拉取」与「用户自主控制」的预期冲突。
安全保证:
- 已存在则完全跳过 (绝不覆盖用户数据)
- 只写 config.json, 失败只记 warning 不阻断启动
"""
store = ExtConfigStore(data_dir)
for config in _presets():
existing = store.get(config.id)
if existing is not None:
# 用户已有此表 (老用户 / 自己重建过) → 一律不动
continue
try:
store.upsert(config)
logger.info("内置扩展表 %s 配置已就绪 (待用户手动获取数据)", config.id)
except Exception as e: # noqa: BLE001
logger.warning("内置扩展表 %s 配置写入失败 (不影响启动): %s", config.id, e)
async def fetch_preset(config_id: str, data_dir: Path) -> int:
"""手动触发某个预设的数据拉取 (供 API 调用)。
Raises:
ValueError: config_id 不是内置预设
Exception: 网络请求/解析/写入失败 (由 API 层转 HTTP 错误)
"""
config = get_preset(config_id)
if config is None:
raise ValueError(f"未知的内置预设: {config_id}")
flatten = _flatten_concept_rows if config_id == "ext_gn_ths" else _flatten_industry_rows
# 确保 config.json 存在 (用户可能从未启动过 ensure_builtin_presets)
store = ExtConfigStore(data_dir)
if store.get(config_id) is None:
store.upsert(config)
n = await _seed_one(config, flatten, data_dir)
logger.info("内置扩展表 %s 手动拉取成功: %d", config_id, n)
return n
+339
View File
@@ -0,0 +1,339 @@
"""扩展数据定时拉取引擎 — 从外部 API 拉取数据写入 Parquet。"""
from __future__ import annotations
import asyncio
import json
import logging
import threading
from datetime import date, datetime, timezone
from functools import reduce
from typing import Any
import httpx
from app.services.ext_data import (
ExtConfig,
ExtConfigStore,
PullConfig,
rows_to_parquet,
)
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# 响应解析
# ---------------------------------------------------------------------------
def _extract_rows(data: Any, path: str) -> list[dict]:
"""按 dot-path 从 JSON 响应中提取行数组。
例: path="data.list" → response["data"]["list"]
如果 path 为空,直接将 data 视为数组。
"""
if not path:
if isinstance(data, list):
return data
raise ValueError("response_path 为空但响应不是数组")
keys = path.split(".")
current = data
for key in keys:
if isinstance(current, dict):
if key not in current:
raise ValueError(f"响应中不存在路径 '{path}',缺失键 '{key}'")
current = current[key]
elif isinstance(current, list):
try:
current = current[int(key)]
except (ValueError, IndexError) as e:
raise ValueError(f"响应路径 '{path}' 解析失败: {e}") from e
else:
raise ValueError(f"响应路径 '{path}' 中间值不是 dict/list: {type(current)}")
if not isinstance(current, list):
raise ValueError(f"路径 '{path}' 指向的不是数组,而是 {type(current)}")
return current
def _apply_field_map(rows: list[dict], field_map: dict[str, str]) -> list[dict]:
"""将外部字段名映射为内部配置字段名。field_map: {外部名: 内部名}。"""
if not field_map:
return rows
mapped = []
for row in rows:
new_row: dict = {}
for k, v in row.items():
mapped_key = field_map.get(k, k)
new_row[mapped_key] = v
mapped.append(new_row)
return mapped
# ---------------------------------------------------------------------------
# 拉取执行
# ---------------------------------------------------------------------------
def _apply_preset_flatten(config_id: str, rows: list[dict]) -> list[dict]:
"""对内置预设 (概念/行业) 应用结构转换, 与 fetch_preset 保持一致。
延迟导入避免与 ext_presets 形成循环依赖。
非预设 id 原样返回。
"""
if config_id not in ("ext_gn_ths", "ext_hy_ths"):
return rows
from app.services.ext_presets import _flatten_concept_rows, _flatten_industry_rows
flatten = _flatten_concept_rows if config_id == "ext_gn_ths" else _flatten_industry_rows
return flatten(rows)
async def fetch_and_ingest(
config: ExtConfig,
data_dir,
) -> tuple[int, str]:
"""执行一次拉取: 请求外部 API → 解析响应 → 写入 Parquet。
Returns:
(rows_written, date_str)
"""
pull = config.pull
if not pull or not pull.url:
raise ValueError("拉取未配置或 URL 为空")
async with httpx.AsyncClient(timeout=30) as client:
headers = pull.headers or {}
kwargs: dict[str, Any] = {"headers": headers}
if pull.method.upper() == "POST" and pull.body:
kwargs["content"] = pull.body
if "content-type" not in {k.lower() for k in headers}:
kwargs["headers"]["Content-Type"] = "application/json"
resp = await client.request(pull.method.upper(), pull.url, **kwargs)
resp.raise_for_status()
# 解析 JSON
try:
data = resp.json()
except Exception as e:
raise ValueError(f"响应不是有效 JSON: {e}") from e
# 提取行
rows = _extract_rows(data, pull.response_path)
if not rows:
raise ValueError("提取到的行数为 0")
# 内置预设 (概念/行业): 应用结构转换, 让产出 schema 与分析页一致。
# 否则 raw 接口列 (concepts/industries 数组、name) 会直接覆盖正确的 part.parquet,
# 导致分析页因找不到维度字段 (所属概念/所属同花顺行业) 而"数据消失"。
# 见 ext_presets._flatten_* —— 手动拉取 / 定时拉取都必须走同一套转换。
rows = _apply_preset_flatten(config.id, rows)
# 字段映射
rows = _apply_field_map(rows, pull.field_map)
# 校验 symbol 列
if rows and "symbol" not in rows[0]:
raise ValueError("数据行中缺少 symbol 字段,请配置 field_map 映射")
# 写入
snap = date.today()
n = rows_to_parquet(rows, config, data_dir, snapshot_date=snap)
return n, snap.isoformat()
# ---------------------------------------------------------------------------
# 调度器
# ---------------------------------------------------------------------------
class PullScheduler:
"""后台调度器:为每个启用了 pull 的 ExtConfig 维护定时任务。
线程安全说明:
refresh()/stop() 可能从主事件循环 (lifespan startup) 或同步路由的
worker 线程 (configure_pull 是 def 而非 async def, FastAPI 丢进线程池)
调用。worker 线程里没有 running loop, 直接 asyncio.create_task 会抛
"no running event loop"。因此对 task 的增删一律通过
call_soon_threadsafe 提交到主循环执行 —— 同一套代码两种调用场景都安全。
"""
def __init__(self) -> None:
self._tasks: dict[str, asyncio.Task] = {}
self._running = False
self._loop: asyncio.AbstractEventLoop | None = None
def start(self, data_dir) -> None:
"""启动调度(在 lifespan startup 调用,主事件循环内)。"""
self._running = True
self._data_dir = data_dir
try:
self._loop = asyncio.get_running_loop()
except RuntimeError:
self._loop = None
logger.info("PullScheduler started")
def _submit(self, fn, *args) -> None:
"""把一个 callable 提交到主事件循环执行 (线程安全)。
startup 在主循环内调用时 fn 立即排队; worker 线程调用时跨线程排队。
两者都通过 call_soon_threadsafe, 保证 _tasks 字典的读写只在主循环里发生。
"""
loop = self._loop
if loop is None or loop.is_closed():
raise RuntimeError(
"PullScheduler: 事件循环不可用 (start() 未在事件循环中调用?)"
)
loop.call_soon_threadsafe(fn, *args)
def stop(self) -> None:
"""停止所有任务 (从 shutdown 调用)。"""
self._running = False
for task in self._tasks.values():
task.cancel()
self._tasks.clear()
logger.info("PullScheduler stopped")
def refresh(self, data_dir) -> None:
"""重新加载配置,更新调度任务(增/删/改)。线程安全。"""
self._data_dir = data_dir
store = ExtConfigStore(data_dir)
configs = store.load_all()
active_ids: set[str] = set()
new_configs: list[ExtConfig] = []
for config in configs:
if not config.pull or not config.pull.enabled or not config.pull.url:
continue
active_ids.add(config.id)
if config.id not in self._tasks:
new_configs.append(config)
# 需要移除的 id (快照当前 task 字典的键, 避免遍历时改字典)
remove_ids = [cid for cid in list(self._tasks) if cid not in active_ids]
# 所有对 _tasks 的修改都提交到主循环里执行, 保证线程安全
def _apply() -> None:
for config in new_configs:
if config.id not in self._tasks: # 二次校验, 防重复
self._tasks[config.id] = self._loop.create_task(
self._run_loop(config)
)
logger.info(
"PullScheduler: scheduled %s (every %d min)",
config.id, config.pull.schedule_minutes,
)
for cid in remove_ids:
task = self._tasks.pop(cid, None)
if task is not None:
task.cancel()
logger.info("PullScheduler: removed %s", cid)
self._submit(_apply)
async def _run_loop(self, config: ExtConfig) -> None:
"""单个配置的定时拉取循环。
策略: 启用后立即执行一次, 之后按 interval 循环。
每次循环重读最新配置 (fresh), interval 取自 fresh.pull.schedule_minutes,
这样用户中途修改间隔也能立即生效 (无需重启)。
"""
try:
while self._running:
# 每轮重读最新配置 — 用户可能修改了 url / interval / enabled
store = ExtConfigStore(self._data_dir)
fresh = store.get(config.id)
if not fresh or not fresh.pull or not fresh.pull.enabled:
break
pull = fresh.pull
# 先执行一次 (启用即拉取, 让用户立刻看到生效)
try:
n, d = await fetch_and_ingest(fresh, self._data_dir)
fresh.pull.last_run = datetime.now(timezone.utc).isoformat()
fresh.pull.last_status = "success"
fresh.pull.last_message = f"{n} rows @ {d}"
fresh.pull.last_rows = n
store.upsert(fresh)
logger.info("PullScheduler: %s success, %d rows", config.id, n)
except Exception as e:
fresh2 = store.get(config.id)
if fresh2 and fresh2.pull:
fresh2.pull.last_run = datetime.now(timezone.utc).isoformat()
fresh2.pull.last_status = "error"
fresh2.pull.last_message = str(e)[:200]
store.upsert(fresh2)
logger.warning("PullScheduler: %s error: %s", config.id, e)
# 间隔取自最新配置 (每次重新读取, 修复改间隔不生效)
interval = max(pull.schedule_minutes * 60, 60) # 至少 60s
# 预告下次运行时间, 供前端展示
next_dt = datetime.now(timezone.utc).timestamp() + interval
latest = store.get(config.id)
if latest and latest.pull:
latest.pull.next_run = datetime.fromtimestamp(
next_dt, tz=timezone.utc
).isoformat()
store.upsert(latest)
await asyncio.sleep(interval)
if not self._running:
break
except asyncio.CancelledError:
pass
async def _run_loop(self, config: ExtConfig) -> None:
"""单个配置的定时拉取循环。
策略: 启用后立即执行一次, 之后按 interval 循环。
每次循环重读最新配置 (fresh), interval 取自 fresh.pull.schedule_minutes,
这样用户中途修改间隔也能立即生效 (无需重启)。
"""
try:
while self._running:
# 每轮重读最新配置 — 用户可能修改了 url / interval / enabled
store = ExtConfigStore(self._data_dir)
fresh = store.get(config.id)
if not fresh or not fresh.pull or not fresh.pull.enabled:
break
pull = fresh.pull
# 先执行一次 (启用即拉取, 让用户立刻看到生效)
try:
n, d = await fetch_and_ingest(fresh, self._data_dir)
fresh.pull.last_run = datetime.now(timezone.utc).isoformat()
fresh.pull.last_status = "success"
fresh.pull.last_message = f"{n} rows @ {d}"
fresh.pull.last_rows = n
store.upsert(fresh)
logger.info("PullScheduler: %s success, %d rows", config.id, n)
except Exception as e:
fresh2 = store.get(config.id)
if fresh2 and fresh2.pull:
fresh2.pull.last_run = datetime.now(timezone.utc).isoformat()
fresh2.pull.last_status = "error"
fresh2.pull.last_message = str(e)[:200]
store.upsert(fresh2)
logger.warning("PullScheduler: %s error: %s", config.id, e)
# 间隔取自最新配置 (每次重新读取, 修复改间隔不生效)
interval = max(pull.schedule_minutes * 60, 60) # 至少 60s
# 预告下次运行时间, 供前端展示
next_dt = datetime.now(timezone.utc).timestamp() + interval
latest = store.get(config.id)
if latest and latest.pull:
latest.pull.next_run = datetime.fromtimestamp(
next_dt, tz=timezone.utc
).isoformat()
store.upsert(latest)
await asyncio.sleep(interval)
if not self._running:
break
except asyncio.CancelledError:
pass
# 全局单例
pull_scheduler = PullScheduler()
@@ -0,0 +1,224 @@
"""向前扩展历史数据 — 完全独立于 daily_pipeline 的盘后管道。
用户从日 K 卡片手动触发,指定往前补的时长 (x 天/月/年)。
流程:
1. 获取当前最早日期
2. 向前拉日 K batch (start = 最早日期 - offset, end = 最早日期)
3. 向前拉除权因子 (同范围)
4. 全量重算 enriched
5. 刷新视图 + 缓存
⚠️ 本模块不导入 daily_pipeline 的任何函数,只复用基础设施:
- kline_sync.sync_and_persist_daily_batch / sync_adj_factor
- indicators.pipeline.run_pipeline
- pipeline_jobs.JobStore
- tickflow.repository.KlineRepository
"""
from __future__ import annotations
import logging
from collections.abc import Callable
from datetime import date, datetime, timedelta
from app.services import kline_sync
from app.services.pipeline_jobs import job_store
from app.tickflow.capabilities import Cap, CapabilitySet
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
def _noop(stage: str, pct: int, msg: str, **kwargs) -> None: # noqa: ARG001
pass
def _invalidate(table: str | None = None) -> None:
from app.api.data import invalidate_data_cache
invalidate_data_cache(table)
def _resolve_universe(capset: CapabilitySet) -> list[str]:
"""解析标的池 — 与 daily_pipeline 独立的副本。"""
if capset.has(Cap.KLINE_DAILY_BATCH):
try:
from app.tickflow.pools import get_pool
all_a = get_pool("CN_Equity_A", refresh=True)
if all_a:
return sorted(all_a)
except Exception as e:
logger.warning("CN_Equity_A pool unavailable: %s", e)
from app.tickflow.pools import DEMO_SYMBOLS, get_pool as _get_pool
from app.config import settings
from pathlib import Path
import polars as pl
base: set[str] = set(DEMO_SYMBOLS)
base.update(_get_pool("watchlist"))
d = Path(settings.data_dir)
inst_path = d / "instruments" / "instruments.parquet"
if inst_path.exists():
try:
inst = pl.read_parquet(inst_path, columns=["symbol"])
base.update(inst["symbol"].to_list())
except Exception as e:
logger.warning("instruments supplement failed: %s", e)
return sorted(base)
def _refresh_single_view(repo: KlineRepository, name: str) -> None:
"""刷新单个 DuckDB 视图。"""
d = repo.store.data_dir.as_posix()
paths = {
"kline_daily": f"{d}/kline_daily/**/*.parquet",
"kline_enriched": f"{d}/kline_daily_enriched/**/*.parquet",
"kline_minute": f"{d}/kline_minute/**/*.parquet",
"adj_factor": f"{d}/adj_factor/**/*.parquet",
"instruments": f"{d}/instruments/**/*.parquet",
}
path = paths.get(name)
if not path:
return
try:
repo.db.execute(
f"CREATE OR REPLACE VIEW {name} AS "
f"SELECT * FROM read_parquet('{path}', union_by_name=true)"
)
except Exception as e:
logger.warning("refresh view %s failed: %s", name, e)
def compute_offset(value: int, unit: str) -> timedelta:
"""将用户输入的 value + unit 转成 timedelta。"""
if unit == "day":
return timedelta(days=value)
elif unit == "month":
return timedelta(days=value * 30)
elif unit == "year":
return timedelta(days=value * 365)
else:
raise ValueError(f"不支持的单位: {unit}")
def run_extend_history(
repo: KlineRepository,
capset: CapabilitySet,
value: int,
unit: str,
on_progress: Callable | None = None,
) -> dict:
"""向前扩展历史数据的主函数。
完全独立于 daily_pipeline.run_now(),不调用其任何逻辑。
返回结果 dict 供 job_store 记录。
"""
emit = on_progress or _noop
# 0. 计算时间偏移
offset = compute_offset(value, unit)
today = date.today()
# 1. 获取当前最早日期
emit("extend_history", 2, "检查当前数据范围…")
earliest = repo.earliest_daily_date()
if not earliest:
return {"error": "本地无日K数据,请先执行一次完整同步"}
new_start = earliest - offset
# 不能超过今天
if new_start >= earliest:
return {"error": "扩展范围无效,请增大时间跨度"}
# 2. 解析标的池
emit("extend_history", 5, "解析标的池…")
universe = _resolve_universe(capset)
if not universe:
return {"error": "标的池为空"}
emit("extend_history", 8, f"标的池: {len(universe)}")
start_str = new_start.strftime("%Y-%m-%d")
end_str = earliest.strftime("%Y-%m-%d")
# 3. 拉日 K
emit("extend_history", 10, f"获取日K [{start_str} ~ {end_str}]…")
logger.info("extend_history: daily K [%s ~ %s], %d symbols", start_str, end_str, len(universe))
def _daily_chunk(cur: int, tot: int) -> None:
emit("extend_history", 10 + int(35 * cur / tot),
f"日K 批次 {cur}/{tot}", stage_pct=int(100 * cur / tot), skip_log=True)
written_daily = kline_sync.sync_and_persist_daily_batch(
universe, repo, capset,
start_date=datetime.combine(new_start, datetime.min.time()),
end_date=datetime.combine(earliest, datetime.min.time()),
on_chunk_done=_daily_chunk,
)
emit("extend_history", 45, f"日K 完成,写入 {written_daily}")
logger.info("extend_history: daily K done, %d rows", written_daily)
_refresh_single_view(repo, "kline_daily")
_invalidate("daily")
# 4. 拉除权因子 (新范围)
written_adj = 0
adj_start = datetime.combine(new_start, datetime.min.time())
adj_end = datetime.combine(today, datetime.min.time())
adj_start_str = new_start.strftime("%Y-%m-%d")
adj_end_str = today.strftime("%Y-%m-%d")
if capset.has(Cap.ADJ_FACTOR):
emit("extend_history", 48, f"获取除权因子 [{adj_start_str} ~ {adj_end_str}]…")
logger.info("extend_history: adj_factor [%s ~ %s]", adj_start_str, adj_end_str)
def _adj_chunk(cur: int, tot: int) -> None:
emit("extend_history", 48 + int(10 * cur / tot),
f"除权因子批次 {cur}/{tot}", stage_pct=int(100 * cur / tot), skip_log=True)
written_adj, _affected = kline_sync.sync_adj_factor(
universe, repo, capset,
start_time=adj_start, end_time=adj_end,
on_chunk_done=_adj_chunk,
)
emit("extend_history", 60, f"除权因子完成,{written_adj}")
logger.info("extend_history: adj_factor done, %d rows", written_adj)
_refresh_single_view(repo, "adj_factor")
_invalidate("adj_factor")
else:
emit("extend_history", 60, "除权因子跳过(无权限)")
logger.info("extend_history: adj_factor skipped, no ADJ_FACTOR capability")
# 5. 全量重算 enriched
emit("extend_history", 65, "全量计算 enriched…")
logger.info("extend_history: full enriched rebuild start")
from app.indicators.pipeline import run_pipeline
written_enriched = run_pipeline()
enriched_dir = repo.store.data_dir / "kline_daily_enriched"
enriched_days = len(list(enriched_dir.glob("date=*"))) if enriched_dir.exists() else 0
emit("extend_history", 92, f"enriched 完成,覆盖 {enriched_days}")
logger.info("extend_history: enriched done, %d days", enriched_days)
_refresh_single_view(repo, "kline_enriched")
_invalidate("enriched")
# 6. 刷新视图
emit("extend_history", 95, "刷新视图…")
_refresh_single_view(repo, "kline_daily")
_refresh_single_view(repo, "kline_enriched")
_refresh_single_view(repo, "adj_factor")
_invalidate(None)
# 7. 统计结果
daily_dir = repo.store.data_dir / "kline_daily"
daily_days = len(list(daily_dir.glob("date=*"))) if daily_dir.exists() else 0
emit("extend_history", 100, f"完成,已扩展至 {new_start}")
return {
"earliest_before": earliest.isoformat(),
"earliest_after": new_start.isoformat(),
"daily_rows": written_daily,
"daily_days": daily_days,
"adj_factor_rows": written_adj,
"enriched_days": enriched_days,
"universe_size": len(universe),
}
@@ -0,0 +1,187 @@
"""AI 财务分析服务 — 读取个股财务数据 → 构建专业提示词 → 流式调用 LLM。
职责: 拉取单只标的的 4 张财务表 → 转成紧凑 JSON → 拼装 CFA 分析师级系统提示词
→ 流式调用 OpenAI 兼容 API → 逐 chunk 吐给前端。
不知道: HTTP、前端、配置持久化。
"""
from __future__ import annotations
import json
import logging
from pathlib import Path
from typing import AsyncIterator
import polars as pl
from app.services.financial_sync import get_financial_df
logger = logging.getLogger(__name__)
# 最多注入的报告期数(最新 N 期),避免上下文爆炸 / token 浪费
_MAX_PERIODS = 4
def _load_stock_financials(data_dir: Path, symbol: str) -> dict[str, list[dict]]:
"""读取该标的的 4 张财务表,返回 {table: [records...]}(按 period_end 降序,截取最新 N 期)。
数值统一做 NaN/Inf → null 清洗,保证 JSON 序列化不报错。
"""
result: dict[str, list[dict]] = {}
for table in ("metrics", "income", "balance_sheet", "cash_flow"):
df = get_financial_df(data_dir, table)
if df.is_empty():
result[table] = []
continue
df = df.filter(pl.col("symbol") == symbol)
if df.is_empty():
result[table] = []
continue
# 按 period_end 降序,截取最新 N 期
if "period_end" in df.columns:
df = df.sort("period_end", descending=True).head(_MAX_PERIODS)
# 清洗 NaN/Inf,转成 JSON 安全的 dict 列表
rows = []
for rec in df.to_dicts():
clean = {}
for k, v in rec.items():
if k == "symbol":
continue # 不需要重复回传 symbol
if isinstance(v, float):
import math
clean[k] = None if not math.isfinite(v) else v
else:
clean[k] = v
rows.append(clean)
result[table] = rows
return result
def _summarize(fins: dict[str, list[dict]]) -> str:
"""生成一行业务摘要,便于 LLM 快速把握数据全貌(行数/期数)。"""
parts = []
for table in ("metrics", "income", "balance_sheet", "cash_flow"):
rows = fins.get(table, [])
if rows:
periods = [r.get("period_end") for r in rows if r.get("period_end")]
parts.append(f"{table}: {len(rows)}期 ({', '.join(str(p) for p in periods[:3])})")
else:
parts.append(f"{table}: 无数据")
return " · ".join(parts)
# ================================================================
# 系统提示词 —— CFA 分析师级,九维分析框架
# ================================================================
_SYSTEM_PROMPT = """你是一位拥有 15 年 A 股投研经验的资深财务分析师(CFA + CPA),服务于专业机构投资者。你的任务是:基于提供的上市公司财务数据,产出一份**严谨、专业、可直接用于投资决策**的财务分析报告。
## 输出规范
用 **Markdown** 格式输出,严格遵循以下结构。不要输出任何 JSON 或代码块,直接输出 Markdown 正文。
### 1. 📌 核心摘要(1-2 句)
用一句话概括该公司的财务画像:盈利质量、成长动能、财务健康度的最关键判断。结尾用【综合评级:★★★☆☆】给出 1-5 星评级。
### 2. ✅ 亮点(2-3 条)
列出最值得关注的**积极信号**,每条用加粗短语领起,配数据支撑。例如盈利高增、ROE 持续提升、现金流充沛等。
### 3. ⚠️ 风险提示(2-3 条)
客观指出**潜在风险或值得警惕的信号**,例如应收激增、存货堆积、经营现金流与净利润背离、债务攀升等。宁可保守,不要回避。
### 4. 📊 分项诊断
用**表格**呈现各维度的诊断结论,列为「维度 / 关键指标 / 判断」。维度包括:
- **盈利能力**:ROE / ROA / 毛利率 / 净利率
- **成长性**:营收同比 / 净利润同比
- **偿债能力**:资产负债率 / 流动比率(用资产/负债估算)
- **现金流**:经营现金流净额 / 与净利润的匹配度
- **营运效率**:存货周转率等(有数据时)
每个判断给「优秀 / 良好 / 一般 / 偏弱 / 警惕」之一,并一句话说明依据。
### 5. 🎯 综合评估与展望
2-3 段总结:该公司当前的财务状态(优秀/稳健/承压/恶化)、核心驱动力、未来需重点跟踪的指标。**结尾给出"投资参考"**:从纯财务质量角度,该股属于(高质量蓝筹 / 稳健成长 / 周期波动 / 财务承压 / 高风险)中的哪一类。
## 分析准则(务必遵守)
1. **数据说话**:每个判断必须引用具体数值(如"营收同比 +28.5%"),严禁空泛套话
2. **纵向对比**:利用多期数据看趋势(改善/恶化),而非只看单期
3. **交叉验证**:经营现金流 vs 净利润(是否造血)、毛利率 vs 费用率(盈利结构)、负债 vs 资产(杠杆)
4. **行业常识**:对照 A 股常识判断水平(如 ROE>15% 优秀,资产负债率>70% 偏高,毛利率<20% 偏低)
5. **诚实中立**:数据不支持时直言"数据不足,无法判断",绝不编造或过度演绎
6. **简明有力**:避免冗长,用专业投资者能扫读的密度输出,总字数 800-1500 字
## 重要免责
报告末尾附一行:"> ⚠️ 本报告由 AI 基于公开财务数据生成,仅供参考,不构成任何投资建议。"
现在请基于下方数据进行分析。"""
def _build_user_prompt(fins: dict[str, list[dict]], symbol: str, focus: str) -> str:
"""构建用户消息:标的代码 + 数据 JSON + 可选关注点。"""
data_json = json.dumps(fins, ensure_ascii=False, indent=2)
lines = [
f"标的标准代码: {symbol}",
f"数据概览: {_summarize(fins)}",
"",
"以下是该标的最新财务数据(JSON 格式,金额单位为元,比率类指标为百分点):",
"```json",
data_json,
"```",
]
if focus.strip():
lines.extend([
"",
f"本次分析请特别关注: {focus.strip()}",
])
return "\n".join(lines)
async def analyze_financials_stream(
data_dir: Path,
symbol: str,
focus: str = "",
) -> AsyncIterator[str]:
"""流式分析:yield 出每个文本 chunk。
- 启动时先 yield 一条 {"type":"meta",...} 让前端显示数据摘要
- 之后逐 chunk yield {"type":"delta","content":"..."}
- 出错时 yield {"type":"error","message":"..."}
- 结束 yield {"type":"done"}
"""
# 1. 加载数据
fins = _load_stock_financials(data_dir, symbol)
total_rows = sum(len(v) for v in fins.values())
if total_rows == 0:
yield json.dumps({"type": "error", "message": f"标的 {symbol} 暂无任何财务数据,请先同步财务表"}, ensure_ascii=False)
return
# 2. meta
yield json.dumps({
"type": "meta",
"symbol": symbol,
"summary": _summarize(fins),
"periods": total_rows,
}, ensure_ascii=False)
# 3. 调用 LLM 流式
try:
from app.services.ai_provider import stream_ai_text
user_prompt = _build_user_prompt(fins, symbol, focus)
async for delta in stream_ai_text(
[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
],
temperature=0.4,
max_tokens=4000,
):
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
except Exception as e: # noqa: BLE001
logger.exception("AI financial analysis failed for %s: %s", symbol, e)
yield json.dumps({"type": "error", "message": f"AI 分析失败: {e}"}, ensure_ascii=False)
return
yield json.dumps({"type": "done"}, ensure_ascii=False)
@@ -0,0 +1,410 @@
"""财务数据独立同步服务。
解耦于 K-line 管道, 自有调度 + 自有存储。
能力门控: Cap.FINANCIAL (Expert 套餐)
"""
from __future__ import annotations
import asyncio
import logging
import threading
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import polars as pl
from app.tickflow.capabilities import Cap, CapabilitySet
logger = logging.getLogger(__name__)
# 每个 API 请求最多 100 个标的
_BATCH_SIZE = 100
# 4 张财务表
FINANCIAL_TABLES = ("metrics", "income", "balance_sheet", "cash_flow")
# ================================================================
# 同步函数
# ================================================================
def _get_symbols(data_dir: Path) -> list[str]:
"""从 instruments 表获取标的列表。"""
inst_path = data_dir / "instruments" / "instruments.parquet"
if not inst_path.exists():
return []
try:
df = pl.read_parquet(inst_path, columns=["symbol"])
return df["symbol"].to_list()
except Exception as e:
logger.warning("读取 instruments 失败: %s", e)
return []
def _sync_table(
table: str,
symbols: list[str],
data_dir: Path,
capset: CapabilitySet,
latest_only: bool = True,
) -> int:
"""同步单张财务表。返回写入的行数。"""
if not capset.has(Cap.FINANCIAL):
logger.info("sync_%s skipped: no FINANCIAL capability", table)
return 0
if not symbols:
logger.warning("sync_%s skipped: no symbols", table)
return 0
from app.tickflow.client import get_client
tf = get_client()
# 分批拉取
api_method = {
"metrics": tf.financials.metrics,
"income": tf.financials.income,
"balance_sheet": tf.financials.balance_sheet,
"cash_flow": tf.financials.cash_flow,
}[table]
all_records: list[dict] = []
total_batches = (len(symbols) + _BATCH_SIZE - 1) // _BATCH_SIZE
for i in range(0, len(symbols), _BATCH_SIZE):
chunk = symbols[i : i + _BATCH_SIZE]
batch_num = i // _BATCH_SIZE + 1
try:
data = api_method(chunk, latest=latest_only)
# data 格式: { "600519.SH": [record, ...], ... }
if isinstance(data, dict):
for sym, records in data.items():
if isinstance(records, list):
for rec in records:
if isinstance(rec, dict):
rec["symbol"] = sym
all_records.append(rec)
logger.debug("sync_%s batch %d/%d: %d records", table, batch_num, total_batches, len(data) if isinstance(data, dict) else 0)
except Exception as e:
logger.warning("sync_%s batch %d/%d failed: %s", table, batch_num, total_batches, e)
if not all_records:
return 0
df = pl.DataFrame(all_records)
if df.is_empty():
return 0
# 确保 symbol 列存在
if "symbol" not in df.columns:
return 0
# 写入 Parquet (全量覆盖)
out_dir = data_dir / "financials" / table
out_dir.mkdir(parents=True, exist_ok=True)
out_file = out_dir / "part.parquet"
df.write_parquet(out_file)
logger.info("sync_%s done: %d records written", table, len(df))
return len(df)
def sync_metrics(data_dir: Path, capset: CapabilitySet) -> int:
"""同步核心财务指标 (metrics)。"""
symbols = _get_symbols(data_dir)
return _sync_table("metrics", symbols, data_dir, capset, latest_only=True)
def sync_income(data_dir: Path, capset: CapabilitySet) -> int:
"""同步利润表。"""
symbols = _get_symbols(data_dir)
return _sync_table("income", symbols, data_dir, capset, latest_only=True)
def sync_balance_sheet(data_dir: Path, capset: CapabilitySet) -> int:
"""同步资产负债表。"""
symbols = _get_symbols(data_dir)
return _sync_table("balance_sheet", symbols, data_dir, capset, latest_only=True)
def sync_cash_flow(data_dir: Path, capset: CapabilitySet) -> int:
"""同步现金流量表。"""
symbols = _get_symbols(data_dir)
return _sync_table("cash_flow", symbols, data_dir, capset, latest_only=True)
def sync_all(data_dir: Path, capset: CapabilitySet) -> dict[str, int]:
"""同步所有财务表。返回 {table: rows}。"""
if not capset.has(Cap.FINANCIAL):
logger.info("sync_all financials skipped: no FINANCIAL capability")
return {}
symbols = _get_symbols(data_dir)
results: dict[str, int] = {}
for table in FINANCIAL_TABLES:
results[table] = _sync_table(table, symbols, data_dir, capset, latest_only=True)
# 同步完成后注册 DuckDB 视图
_refresh_financials_views(data_dir)
return results
# ================================================================
# DuckDB 视图
# ================================================================
def _refresh_financials_views(data_dir: Path) -> None:
"""刷新财务表 DuckDB 视图 (在 DataStore.db 上注册)。"""
d = data_dir.as_posix()
views = {
"financials_metrics": f"{d}/financials/metrics/*.parquet",
"financials_income": f"{d}/financials/income/*.parquet",
"financials_balance_sheet": f"{d}/financials/balance_sheet/*.parquet",
"financials_cash_flow": f"{d}/financials/cash_flow/*.parquet",
}
for name, path in views.items():
out = data_dir / "financials" / name.replace("financials_", "") / "part.parquet"
if not out.exists():
continue
# 视图注册需要由 DataStore 完成,这里只做日志
logger.debug("financial parquet ready: %s (%d rows)", name, out.stat().st_size)
def get_financial_df(data_dir: Path, table: str) -> pl.DataFrame:
"""读取本地财务 Parquet。"""
path = data_dir / "financials" / table / "part.parquet"
if not path.exists():
return pl.DataFrame()
try:
return pl.read_parquet(path)
except Exception as e:
logger.warning("读取 financials/%s 失败: %s", table, e)
return pl.DataFrame()
# ================================================================
# 调度器
# ================================================================
class FinancialScheduler:
"""独立调度器: 每周同步 metrics, 每季度同步三张报表。"""
def __init__(self) -> None:
self._task: asyncio.Task | None = None
self._running = False
self._data_dir: Path | None = None
self._capset: CapabilitySet | None = None
self._lock = threading.Lock()
self._last_sync: dict[str, str] = {} # {table: iso_timestamp}
# 手动同步(run_now)是否正在进行。前端据此显示"同步中"并防重复点击。
self._is_syncing = False
def start(self, data_dir: Path, capset: CapabilitySet, *, auto_schedule: bool = False) -> None:
"""初始化调度器,并按需启动周期同步后台任务。
auto_schedule=False (默认): 仅初始化 (设置数据目录/能力 + 恢复 last_sync),
供 /api/financials/sync/* 手动同步使用, 不启动自动调度。
auto_schedule=True: 额外启动每周一次的 metrics 自动同步 (启动后 60s 首跑)。
"""
# 先记录 data_dir/capset, 即使当前无 FINANCIAL 也保留引用:
# 用户稍后在「设置」页升级到 Expert Key 时, update_capabilities() 会把新 capset
# 推进来,trigger()/run_now() 才能用上 FINANCIAL。否则 _capset 永远是 None,
# 即便 app.state.capabilities 已更新, 调度器仍报 "no FINANCIAL capability"。
self._data_dir = data_dir
self._capset = capset
if not capset.has(Cap.FINANCIAL):
logger.info("FinancialScheduler skipped: no FINANCIAL capability")
return
# 从持久化恢复上次同步时间: 重启后前端仍能显示真实最后同步时间,而非"尚未同步"
try:
from app.services import preferences
restored = dict(preferences.get_financial_sync_times())
# 老用户迁移兜底: 若某表在 preferences 无记录但 parquet 已存在(升级前同步过),
# 用 parquet 文件的修改时间作为同步时间并补写持久化。
for table in FINANCIAL_TABLES:
if table in restored:
continue
parquet = data_dir / "financials" / table / "part.parquet"
if parquet.exists():
mtime = datetime.fromtimestamp(parquet.stat().st_mtime, tz=timezone.utc).isoformat()
restored[table] = mtime
preferences.set_financial_sync_time(table, mtime)
logger.info("FinancialScheduler backfilled last_sync for %s from parquet mtime", table)
self._last_sync = restored
if self._last_sync:
logger.info("FinancialScheduler restored last_sync: %s", list(self._last_sync.keys()))
except Exception as e: # noqa: BLE001
logger.warning("restore financial_sync_times failed: %s", e)
if not auto_schedule:
# 仅初始化 (手动同步用), 不启动周期任务。
logger.info("FinancialScheduler initialized (auto-schedule disabled; manual sync only)")
return
self._running = True
self._task = asyncio.create_task(self._run_loop())
logger.info("FinancialScheduler started (auto-schedule enabled)")
def _record_sync(self, table: str) -> None:
"""记录一张表的同步完成时间: 更新内存 + 持久化到 preferences.json。
持久化确保即使重启,前端 /status 仍返回真实的最后同步时间,
不会错误地显示"尚未同步"
"""
ts = datetime.now(timezone.utc).isoformat()
self._last_sync[table] = ts
try:
from app.services import preferences
preferences.set_financial_sync_time(table, ts)
except Exception as e: # noqa: BLE001
logger.warning("persist financial_sync_time(%s) failed: %s", e)
def update_capabilities(self, capset: CapabilitySet) -> None:
"""刷新调度器持有的能力集。
用户在「设置」页新增/清除 API Key 后, settings API 会重新探测能力并更新
app.state.capabilities; 必须同步推给本调度器, 否则 trigger()/run_now() 仍读
启动时的旧 capset, 即便 app.state 已含 FINANCIAL, 调度器仍报
"no FINANCIAL capability" 而拒绝同步 (表现为前端「全部同步」按钮闪一下无动作)。
"""
prev = self._capset
self._capset = capset
had = bool(prev) and prev.has(Cap.FINANCIAL)
now = capset.has(Cap.FINANCIAL)
if had != now:
logger.info(
"FinancialScheduler capabilities updated: FINANCIAL %s -> %s", had, now
)
def stop(self) -> None:
self._running = False
if self._task:
self._task.cancel()
self._task = None
logger.info("FinancialScheduler stopped")
async def _run_loop(self) -> None:
"""每周执行一次 metrics 同步。"""
try:
while self._running:
# 首次启动等 60s, 之后每 7 天执行一次
await asyncio.sleep(60)
if not self._running:
break
# 每周: 只同步 metrics
try:
rows = sync_metrics(self._data_dir, self._capset)
self._record_sync("metrics")
logger.info("FinancialScheduler: metrics synced, %d rows", rows)
except Exception as e:
logger.warning("FinancialScheduler: metrics sync failed: %s", e)
# 等待下一次 (7天)
for _ in range(7 * 24 * 60): # 每分钟检查一次 _running
if not self._running:
break
await asyncio.sleep(60)
except asyncio.CancelledError:
pass
def _run_body(self, table: str | None) -> dict[str, int]:
"""同步逻辑本体(不加锁,假设调用方已持有 _is_syncing)。
table=None 同步全部 4 张表;否则只同步指定表。
每张表完成立即更新 last_sync,让前端轮询 /status 能看到进度递增。
"""
if table:
fn = {
"metrics": sync_metrics,
"income": sync_income,
"balance_sheet": sync_balance_sheet,
"cash_flow": sync_cash_flow,
}.get(table)
if not fn:
return {}
rows = fn(self._data_dir, self._capset)
self._record_sync(table)
return {table: rows}
# 全部同步
symbols = _get_symbols(self._data_dir)
result: dict[str, int] = {}
for t in FINANCIAL_TABLES:
result[t] = _sync_table(t, symbols, self._data_dir, self._capset, latest_only=True)
self._record_sync(t)
_refresh_financials_views(self._data_dir)
return result
def run_now(self, table: str | None = None) -> dict[str, int]:
"""同步执行一次同步(阻塞调用线程)。
⚠ 全量同步需数分钟,务必在后台线程调用,不要直接在 HTTP 请求线程里阻塞,
否则请求会长时间 pending 直至被浏览器/代理超时掐断(表现为"点击无反应")。
HTTP 接口应调用 trigger() 立即返回,再让前端轮询 /status.syncing 看进度。
用 _is_syncing 标志防并发:若已有同步在进行,本次直接跳过,
避免重复请求拖慢服务端 / 触发上游限流。
"""
if not self._capset or not self._capset.has(Cap.FINANCIAL):
return {}
with self._lock:
if self._is_syncing:
logger.info("financial sync skipped: already running")
return {"_skipped": 1}
self._is_syncing = True
try:
return self._run_body(table)
finally:
with self._lock:
self._is_syncing = False
def trigger(self, table: str | None = None) -> dict[str, int]:
"""触发一次同步(非阻塞,立即返回)。
在后台线程执行同步体,HTTP 请求无需等待。
返回 {"started": True/False}:
- False = 能力不足或已有同步在进行(被防并发跳过)
- True = 已在后台开始,前端应轮询 /status.syncing 观察进度
⚠ _is_syncing 在此处置 True(持锁),确保 trigger 返回时前端轮询
/status 已能看到 syncing=True,无竞态窗口;同时防止快速重复点击
启动多个后台线程。后台线程复用 _run_body 执行真正的同步逻辑。
"""
if not self._capset or not self._capset.has(Cap.FINANCIAL):
return {"started": False, "reason": "no FINANCIAL capability"}
with self._lock:
if self._is_syncing:
logger.info("financial sync trigger skipped: already running")
return {"started": False, "reason": "already running"}
# 持锁置位:保证 trigger 返回前 syncing 已为 True
self._is_syncing = True
def _bg() -> None:
try:
self._run_body(table)
except Exception as e: # noqa: BLE001
logger.exception("background financial sync failed: %s", e)
finally:
with self._lock:
self._is_syncing = False
t = threading.Thread(target=_bg, name="financial-sync", daemon=True)
t.start()
logger.info("financial sync triggered in background: table=%s", table or "all")
return {"started": True}
@property
def is_syncing(self) -> bool:
"""手动同步是否正在进行(供 /status 返回,前端据此显示"同步中")。"""
with self._lock:
return self._is_syncing
@property
def last_sync(self) -> dict[str, str]:
return dict(self._last_sync)
# 全局单例
financial_scheduler = FinancialScheduler()
+360
View File
@@ -0,0 +1,360 @@
"""指数 / ETF 数据同步服务。
标的列表优先用免费的 exchanges.get_instruments(type=index/etf) 拉取
(None/Free 档均可用,无需 quote.pool 权限);付费档可额外用
quotes.get_by_universes 作为补充来源。日K统一走 klines.batch。
"""
from __future__ import annotations
import logging
import gc
from collections.abc import Callable
from datetime import datetime, timedelta
import polars as pl
from app.indicators.pipeline import compute_enriched
from app.services import kline_sync, preferences
from app.tickflow.capabilities import Cap, CapabilitySet
from app.tickflow.client import get_client
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
# exchanges.get_instruments 查询的交易所(沪深京)
_EXCHANGES = ["SH", "SZ", "BJ"]
def _quotes_to_index_instruments(resp) -> pl.DataFrame:
"""将 TickFlow quotes 响应(get_by_universes)规范为指数 instruments。
付费档(Starter+)的补充来源,免费档用不到。
"""
if resp is None:
return pl.DataFrame()
if isinstance(resp, pl.DataFrame):
df = resp
elif hasattr(resp, "columns"):
df = pl.from_pandas(resp.reset_index() if hasattr(resp, "reset_index") else resp)
else:
rows: list[dict] = []
for q in resp or []:
item = q if isinstance(q, dict) else {}
ext = item.get("ext") or {}
symbol = item.get("symbol")
if not symbol:
continue
rows.append({
"symbol": str(symbol),
"name": ext.get("name") or item.get("name") or str(symbol),
})
df = pl.DataFrame(rows)
if df.is_empty() or "symbol" not in df.columns:
return pl.DataFrame()
rename = {"ts_code": "symbol"}
df = df.rename({k: v for k, v in rename.items() if k in df.columns})
if "name" not in df.columns:
if "ext" in df.columns:
df = df.with_columns(pl.col("symbol").cast(pl.Utf8).alias("name"))
else:
df = df.with_columns(pl.col("symbol").cast(pl.Utf8).alias("name"))
result = df.select([
pl.col("symbol").cast(pl.Utf8),
pl.col("name").cast(pl.Utf8),
]).with_columns([
pl.col("symbol").str.split(".").list.first().alias("code"),
pl.lit("index").alias("asset_type"),
])
return result.unique(subset=["symbol"], keep="last").sort("symbol")
def _fetch_instruments_by_type(instrument_type: str, asset_type_label: str) -> pl.DataFrame:
"""用免费的 exchanges.get_instruments 拉取指定类型的标的列表。
None/Free 档均可使用(标的信息查询免费开放)。
instrument_type: 'index' / 'etf'
asset_type_label: 写入 instruments 表的 asset_type 标记('index' / 'etf')
"""
tf = get_client()
rows: list[dict] = []
for ex in _EXCHANGES:
try:
items = tf.exchanges.get_instruments(ex, instrument_type=instrument_type)
for it in items or []:
item = it if isinstance(it, dict) else {}
symbol = item.get("symbol")
if not symbol:
continue
rows.append({
"symbol": str(symbol),
"name": item.get("name") or str(symbol),
})
except Exception as e: # noqa: BLE001
logger.warning("get_instruments(%s, type=%s) failed: %s", ex, instrument_type, e)
if not rows:
return pl.DataFrame()
return (
pl.DataFrame(rows)
.with_columns([
pl.col("symbol").str.split(".").list.first().alias("code"),
pl.lit(asset_type_label).alias("asset_type"),
])
.unique(subset=["symbol"], keep="last")
.sort("symbol")
)
def sync_index_instruments(
repo: KlineRepository,
pull_index: bool = True,
pull_etf: bool = True,
) -> int:
"""同步指数 / ETF 标的维表,返回标的总数。
新版物理分开保存: 指数写 instruments_index, ETF 写 instruments_etf。
读取层仍兼容旧版 instruments_index 中 asset_type='etf' 的历史数据。
"""
index_parts: list[pl.DataFrame] = []
etf_parts: list[pl.DataFrame] = []
# 1) 免费通道:按开关分别拉 index / etf
if pull_index:
index_df = _fetch_instruments_by_type("index", "index")
if not index_df.is_empty():
index_parts.append(index_df)
if pull_etf:
etf_df = _fetch_instruments_by_type("etf", "etf")
if not etf_df.is_empty():
etf_parts.append(etf_df)
# 2) 付费补充:Starter+ 用 get_by_universes 补指数(仅当开启指数拉取)
if pull_index:
capset = None
try:
from app.tickflow import policy
capset = policy.detect_capabilities(force=False)
except Exception: # noqa: BLE001
pass
if capset is not None and capset.has(Cap.QUOTE_POOL):
tf = get_client()
for kwargs in (
{"universes": ["CN_Index"]},
{"universes": ["CN_Index"], "as_dataframe": False},
):
try:
resp = tf.quotes.get_by_universes(**kwargs)
if resp is not None and len(resp) > 0:
sup = _quotes_to_index_instruments(resp)
if not sup.is_empty():
index_parts.append(sup)
break
except Exception as e: # noqa: BLE001
logger.debug("CN_Index universe supplement failed: %s", e)
total = 0
if index_parts:
index_inst = pl.concat(index_parts, how="diagonal_relaxed").unique(subset=["symbol"], keep="last").sort("symbol")
if not index_inst.is_empty():
repo.save_index_instruments(index_inst)
total += index_inst.height
if etf_parts:
etf_inst = pl.concat(etf_parts, how="diagonal_relaxed").unique(subset=["symbol"], keep="last").sort("symbol")
if not etf_inst.is_empty():
repo.save_etf_instruments(etf_inst)
total += etf_inst.height
if total == 0:
logger.warning("指数/ETF 标的列表为空(pull_index=%s, pull_etf=%s)", pull_index, pull_etf)
return 0
repo.refresh_index_views()
logger.info("指数/ETF 标的同步完成: %d", total)
return total
def sync_etf_instruments(repo: KlineRepository) -> int:
"""单独同步 ETF 标的维表(返回 ETF 数量)。"""
etf_df = _fetch_instruments_by_type("etf", "etf")
if etf_df.is_empty():
return 0
repo.save_etf_instruments(etf_df)
repo.refresh_index_views()
return etf_df.height
def sync_and_persist_index_daily(
repo: KlineRepository,
capset: CapabilitySet,
count: int | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
symbols_override: list[str] | None = None,
on_chunk_done: Callable[[int, int], None] | None = None,
) -> int:
"""同步指数/ETF 日K到独立 parquet,并计算 enriched。
symbols_override 非空时,只拉这些代码(跳过 instruments 表),用于自定义范围。
否则取 index_instruments 表全量(指数+ETF 合并存储)。
on_chunk_done(current, total) 每个批次完成后回调。
"""
if not capset.has(Cap.KLINE_DAILY_BATCH):
return 0
if symbols_override:
symbols = sorted(set(s for s in symbols_override if s))
if not symbols:
return 0
else:
instruments = repo.get_index_instruments()
if instruments.is_empty():
sync_index_instruments(repo, pull_index=True, pull_etf=False)
instruments = repo.get_index_instruments()
if not instruments.is_empty() and "asset_type" in instruments.columns:
instruments = instruments.filter(pl.col("asset_type") != "etf")
if instruments.is_empty() or "symbol" not in instruments.columns:
return 0
symbols = sorted(set(instruments["symbol"].to_list()))
lim = capset.limits(Cap.KLINE_DAILY_BATCH)
batch_size = preferences.get_index_daily_batch_size()
if lim and lim.batch:
batch_size = min(batch_size, lim.batch)
rpm = lim.rpm if lim else None
end_time = end_date or datetime.now()
start_time = start_date or (end_time - timedelta(days=365))
total_rows = 0
interval = (60.0 / rpm) if rpm else 0
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
for i, chunk in enumerate(chunks):
if i > 0 and interval > 0 and len(chunks) > rpm:
import time
time.sleep(interval)
raw = kline_sync.sync_daily_batch(
chunk,
count=count,
batch_size=None,
start_time=start_time,
end_time=end_time,
)
if raw.is_empty():
continue
repo.append_index_daily(raw)
enriched = compute_enriched(raw, factors=None, instruments=None)
repo.append_index_enriched(enriched)
total_rows += raw.height
logger.info("index/etf daily synced: %d/%d chunks, +%d rows", i + 1, len(chunks), raw.height)
if on_chunk_done:
on_chunk_done(i + 1, len(chunks))
del raw, enriched
gc.collect()
repo.refresh_index_views()
return total_rows
def _load_etf_factors(repo: KlineRepository) -> pl.DataFrame:
factor_path = repo.store.data_dir / "adj_factor_etf" / "all.parquet"
if not factor_path.exists():
return pl.DataFrame()
try:
return pl.read_parquet(factor_path)
except Exception as e: # noqa: BLE001
logger.warning("ETF 复权因子读取失败: %s", e)
return pl.DataFrame()
def sync_etf_adj_factor(
symbols: list[str],
repo: KlineRepository,
capset: CapabilitySet,
start_time: datetime | None = None,
end_time: datetime | None = None,
on_chunk_done=None,
) -> tuple[int, list[str]]:
"""同步 ETF 复权因子;失败由调用方降级为 warning。"""
return kline_sync.sync_adj_factor(
symbols,
repo,
capset,
start_time=start_time,
end_time=end_time,
on_chunk_done=on_chunk_done,
asset_type="etf",
)
def sync_and_persist_etf_daily(
repo: KlineRepository,
capset: CapabilitySet,
count: int | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
symbols_override: list[str] | None = None,
on_chunk_done: Callable[[int, int], None] | None = None,
) -> int:
"""同步 ETF 日K到独立 kline_etf_* parquet,并计算 ETF enriched。
on_chunk_done(current, total) 每个批次完成后回调。
"""
if not capset.has(Cap.KLINE_DAILY_BATCH):
return 0
if symbols_override:
symbols = sorted(set(s for s in symbols_override if s))
else:
instruments = repo.get_etf_instruments()
if instruments.is_empty():
sync_etf_instruments(repo)
instruments = repo.get_etf_instruments()
if instruments.is_empty() or "symbol" not in instruments.columns:
return 0
symbols = sorted(set(instruments["symbol"].to_list()))
if not symbols:
return 0
lim = capset.limits(Cap.KLINE_DAILY_BATCH)
batch_size = preferences.get_index_daily_batch_size()
if lim and lim.batch:
batch_size = min(batch_size, lim.batch)
rpm = lim.rpm if lim else None
end_time = end_date or datetime.now()
start_time = start_date or (end_time - timedelta(days=365))
total_rows = 0
interval = (60.0 / rpm) if rpm else 0
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
factors = _load_etf_factors(repo)
for i, chunk in enumerate(chunks):
if i > 0 and interval > 0 and len(chunks) > rpm:
import time
time.sleep(interval)
raw = kline_sync.sync_daily_batch(
chunk,
count=count,
batch_size=None,
start_time=start_time,
end_time=end_time,
)
if raw.is_empty():
continue
repo.append_etf_daily(raw)
batch_factors = factors.filter(pl.col("symbol").is_in(chunk)) if not factors.is_empty() else factors
# ETF 使用复权和通用技术指标;不传 instruments,避免套用 A股涨跌停/连板逻辑。
enriched = compute_enriched(raw, factors=batch_factors, instruments=None)
repo.append_etf_enriched(enriched)
total_rows += raw.height
logger.info("etf daily synced: %d/%d chunks, +%d rows", i + 1, len(chunks), raw.height)
if on_chunk_done:
on_chunk_done(i + 1, len(chunks))
del raw, enriched
gc.collect()
repo.refresh_index_views()
return total_rows
@@ -0,0 +1,122 @@
"""标的维表同步服务。
盘前 9:10 调用 tf.exchanges.get_instruments("SH"/"SZ"/"BJ", type="stock")
获取全量标的元数据,flatten ext 字段,写入 instruments.parquet。
Starter+ 盘后可用 quotes.get(universes) 顺便补充 name。
"""
from __future__ import annotations
import logging
from datetime import date
from pathlib import Path
import polars as pl
from app.tickflow.client import get_client
logger = logging.getLogger(__name__)
_EXCHANGES = ["SH", "SZ", "BJ"]
def _flatten_instruments(items: list[dict]) -> list[dict]:
"""把 SDK 返回的 Instrument 列表 flatten 成扁平行。"""
rows = []
for item in items:
row = {
"symbol": item.get("symbol"),
"name": item.get("name"),
"code": item.get("code"),
"exchange": item.get("exchange"),
"region": item.get("region"),
"type": item.get("type"),
}
ext = item.get("ext") or {}
row["listing_date"] = ext.get("listing_date")
row["total_shares"] = ext.get("total_shares")
row["float_shares"] = ext.get("float_shares")
row["tick_size"] = ext.get("tick_size")
row["limit_up"] = ext.get("limit_up")
row["limit_down"] = ext.get("limit_down")
rows.append(row)
return rows
def sync_instruments(data_dir: Path) -> int:
"""全量同步标的维表 → data/instruments/instruments.parquet。
返回写入的行数。
"""
tf = get_client()
all_rows: list[dict] = []
for ex in _EXCHANGES:
try:
items = tf.exchanges.get_instruments(ex, instrument_type="stock")
if items:
all_rows.extend(_flatten_instruments(items))
logger.info("instruments %s: %d stocks", ex, len(items))
except Exception as e:
logger.warning("get_instruments(%s) failed: %s", ex, e)
if not all_rows:
return 0
df = pl.DataFrame(all_rows)
df = df.with_columns(pl.lit(date.today()).alias("as_of"))
out = data_dir / "instruments" / "instruments.parquet"
out.parent.mkdir(parents=True, exist_ok=True)
df.write_parquet(out)
logger.info("instruments synced: %d rows → %s", df.height, out)
return df.height
def enrich_names_from_quotes(
data_dir: Path,
quotes_data: list[dict],
) -> int:
"""从 quotes 响应中提取 name,更新 instruments 维表(兜底补充)。
盘后 quotes.get(universes) 返回的数据中包含 ext.name
用来补充 instruments 中可能缺失的 name。
"""
if not quotes_data:
return 0
# 构建 symbol → name 映射
name_map: dict[str, str] = {}
for q in quotes_data:
symbol = q.get("symbol", "")
ext = q.get("ext") or {}
name = ext.get("name") or q.get("name", "")
if symbol and name:
name_map[symbol] = name
if not name_map:
return 0
inst_path = data_dir / "instruments" / "instruments.parquet"
if not inst_path.exists():
return 0
df = pl.read_parquet(inst_path)
# 只更新空 name 的行
updates = pl.DataFrame({
"symbol": list(name_map.keys()),
"_new_name": list(name_map.values()),
})
df = df.join(updates, on="symbol", how="left")
df = df.with_columns(
pl.when(pl.col("name").is_null() | (pl.col("name") == ""))
.then(pl.col("_new_name"))
.otherwise(pl.col("name"))
.alias("name"),
).drop("_new_name")
df.write_parquet(inst_path)
logger.info("instruments name enriched from quotes: %d names", len(name_map))
return len(name_map)
+674
View File
@@ -0,0 +1,674 @@
"""日 K 同步服务(§7.7 Step 1)。
调度器在 capability 允许下,把符号集合的日 K 批量同步到本地 Parquet。
策略:
- 日 K 仅使用 `kline.daily.batch`
- 除权因子仅使用 `adj_factor`
"""
from __future__ import annotations
import logging
import time
from collections.abc import Callable
from datetime import datetime, timedelta
import polars as pl
from app.indicators.pipeline import filter_halt_days
from app.tickflow.capabilities import Cap, CapabilitySet
from app.tickflow.client import get_client
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
# 标准列(无论 SDK 返回什么形状,我们把它规范成这套)
CANONICAL_DAILY_COLS = [
"symbol", "date", "open", "high", "low", "close", "volume", "amount",
]
def _normalize_daily(df_in, default_symbol: str | None = None) -> pl.DataFrame:
"""把 SDK 返回的 pandas/任意 DataFrame 规范成 canonical 列。"""
if df_in is None or len(df_in) == 0:
return pl.DataFrame()
if not isinstance(df_in, pl.DataFrame):
df = pl.from_pandas(df_in.reset_index() if hasattr(df_in, "reset_index") else df_in)
else:
df = df_in
# 兼容字段名差异
rename_map = {
"ts_code": "symbol",
"trade_date": "date",
"vol": "volume",
"amt": "amount",
"datetime": "date",
}
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
if "symbol" not in df.columns and default_symbol is not None:
df = df.with_columns(pl.lit(default_symbol).alias("symbol"))
# 类型规范
if "date" in df.columns and df.schema["date"] != pl.Date:
df = df.with_columns(pl.col("date").cast(pl.Date, strict=False))
for col in ("open", "high", "low", "close"):
if col in df.columns:
df = df.with_columns(pl.col(col).cast(pl.Float64, strict=False))
for col in ("volume", "amount"):
if col in df.columns:
df = df.with_columns(pl.col(col).cast(pl.Float64, strict=False))
# 过滤停牌日 (open/high 为 0; close 可能被填充为前收盘价, 不能用全零判断)
df = filter_halt_days(df)
# 只保留 canonical 列
keep = [c for c in CANONICAL_DAILY_COLS if c in df.columns]
return df.select(keep)
def sync_daily_batch(symbols: list[str],
count: int | None = None,
batch_size: int | None = None,
rpm: int | None = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
on_chunk_done: Callable[[int, int], None] | None = None) -> pl.DataFrame:
"""批量拉取多股日 K。
优先使用 start_time / end_time 区间 + count=10000,确保覆盖完整时间段。
仅传 count 时按条数回溯。
"""
tf = get_client()
out: list[pl.DataFrame] = []
interval = (60.0 / rpm) if rpm else 0
if batch_size is None:
chunks = [symbols]
else:
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
for i, chunk in enumerate(chunks):
if i > 0 and interval > 0 and len(chunks) > rpm:
time.sleep(interval)
try:
if start_time and end_time:
raw = tf.klines.batch(
chunk, period="1d", adjust="none",
start_time=_datetime_to_ms(start_time),
end_time=_datetime_to_ms(end_time),
count=10000,
as_dataframe=True, show_progress=False,
)
else:
raw = tf.klines.batch(chunk, period="1d", count=count or 250, adjust="none",
as_dataframe=True, show_progress=False)
except Exception as e: # noqa: BLE001
logger.warning("batch fetch failed for %d symbols: %s", len(chunk), e)
continue
# 兼容两种形态:dict[sym → df] 和扁平 df
if isinstance(raw, dict):
for sym, sub in raw.items():
if sub is None or len(sub) == 0:
continue
out.append(_normalize_daily(sub, default_symbol=sym))
elif raw is not None and len(raw) > 0:
out.append(_normalize_daily(raw))
if on_chunk_done:
on_chunk_done(i + 1, len(chunks))
if not out:
return pl.DataFrame()
return pl.concat(out, how="diagonal_relaxed")
def sync_and_persist_daily_batch(
symbols: list[str],
repo: KlineRepository,
capset: CapabilitySet,
count: int | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
on_chunk_done: Callable[[int, int], None] | None = None,
) -> int:
"""批量同步日 K 并落到 Parquet。返回写入的行数。
start_date/end_date: 外部传入的时间范围(由 pipeline 根据已有数据计算)。
未传入时默认拉最近 1 年。
"""
if not symbols or not capset.has(Cap.KLINE_DAILY_BATCH):
return 0
lim = capset.limits(Cap.KLINE_DAILY_BATCH)
batch_size = lim.batch if lim and lim.batch else 100
rpm = lim.rpm if lim else None
end_time = end_date or datetime.now()
start_time = start_date or (end_time - timedelta(days=365))
df = sync_daily_batch(
symbols, count=count, batch_size=batch_size, rpm=rpm,
start_time=start_time, end_time=end_time,
on_chunk_done=on_chunk_done,
)
if df.is_empty():
return 0
repo.append_daily(df)
try:
d = repo.store.data_dir.as_posix()
repo.db.execute(
f"""CREATE OR REPLACE VIEW kline_daily AS
SELECT * FROM read_parquet('{d}/kline_daily/**/*.parquet', union_by_name=true)"""
)
except Exception as e: # noqa: BLE001
logger.warning("refresh view failed: %s", e)
return df.height
def sync_daily_by_quotes(repo: KlineRepository) -> int:
"""用实时行情接口拉全市场当日数据,覆写 kline_daily 今天分区。
一个请求覆盖 ~5500 只股票,比 batch K-line 快几个数量级。
返回写入的行数。
"""
from datetime import date as _date
from app.tickflow.client import get_client
tf = get_client()
try:
resp = tf.quotes.get_by_universes(universes=["CN_Equity_A"])
except Exception as e:
logger.warning("get_by_universes failed: %s", e)
return 0
if not resp:
logger.warning("get_by_universes returned empty")
return 0
records = []
for q in resp:
ext = q.get("ext") or {}
records.append({
"symbol": q.get("symbol"),
"open": q.get("open"),
"high": q.get("high"),
"low": q.get("low"),
"close": q.get("last_price"),
"volume": q.get("volume"),
"amount": q.get("amount"),
})
df = pl.DataFrame(records)
if df.is_empty():
return 0
today = _date.today()
daily_df = df.with_columns(pl.lit(today).cast(pl.Date).alias("date"))
# 过滤停牌 (open/high 为 0; close 可能被填充为前收盘价, 不能用全零判断)
daily_df = filter_halt_days(daily_df)
repo.flush_live_daily(daily_df)
logger.info("sync_daily_by_quotes: %d symbols flushed for %s", daily_df.height, today)
return daily_df.height
def _normalize_adj_factor(raw) -> pl.DataFrame:
"""Normalize SDK ex_factors response to symbol/trade_date/ex_factor."""
if raw is None or len(raw) == 0:
return pl.DataFrame()
if isinstance(raw, dict):
rows: list[dict] = []
for sym, values in raw.items():
for item in values or []:
row = dict(item or {})
row.setdefault("symbol", sym)
rows.append(row)
df = pl.DataFrame(rows) if rows else pl.DataFrame()
elif isinstance(raw, pl.DataFrame):
df = raw
else:
df = pl.from_pandas(raw.reset_index() if hasattr(raw, "reset_index") else raw)
if df.is_empty():
return df
# rename: timestamp/date → trade_date, adj_factor → ex_factor
# 注意: 新版 SDK 可能同时返回 timestamp 和 trade_date (或 adj_factor 和 ex_factor),
# 直接 rename 会产生重复列报错。仅当目标列不存在时才 rename。
rename_map: dict[str, str] = {}
for src, dst in (("timestamp", "trade_date"), ("date", "trade_date"), ("adj_factor", "ex_factor")):
if src in df.columns and dst not in df.columns:
rename_map[src] = dst
df = df.rename(rename_map)
if "trade_date" in df.columns:
if df.schema["trade_date"] in {pl.Int64, pl.Int32, pl.UInt64, pl.UInt32, pl.Float64, pl.Float32}:
df = df.with_columns(
pl.from_epoch(pl.col("trade_date").cast(pl.Int64), time_unit="ms").dt.date().alias("trade_date")
)
else:
df = df.with_columns(pl.col("trade_date").cast(pl.Date, strict=False))
if "ex_factor" in df.columns:
df = df.with_columns(pl.col("ex_factor").cast(pl.Float64, strict=False))
cols = [c for c in ["symbol", "trade_date", "ex_factor"] if c in df.columns]
if len(cols) < 3:
return pl.DataFrame()
return df.select(cols).drop_nulls()
def sync_adj_factor(symbols: list[str], repo: KlineRepository,
capset: CapabilitySet,
start_time: datetime | None = None,
end_time: datetime | None = None,
on_chunk_done: Callable[[int, int], None] | None = None,
asset_type: str = "stock") -> tuple[int, list[str]]:
"""同步除权因子(Starter+)。SDK 接口:`tf.klines.ex_factors(symbols=...)`。
支持增量: 传 start_time/end_time 只拉取该时间范围内的新除权事件。
返回 (写入行数, 受影响的 symbol 列表) — 供 enriched 局部重算使用。
"""
if not capset.has(Cap.ADJ_FACTOR) or not symbols:
return 0, []
tf = get_client()
lim = capset.limits(Cap.ADJ_FACTOR)
batch_size = lim.batch if lim and lim.batch else 50
rpm = lim.rpm if lim else 30
interval = 60.0 / rpm if rpm else 0
# 构建 SDK 参数
sdk_kwargs: dict = {"as_dataframe": True, "batch_size": batch_size, "show_progress": False}
if start_time:
sdk_kwargs["start_time"] = _datetime_to_ms(start_time)
if end_time:
sdk_kwargs["end_time"] = _datetime_to_ms(end_time)
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
all_dfs: list[pl.DataFrame] = []
for i, chunk in enumerate(chunks):
if i > 0 and interval > 0 and len(chunks) > rpm:
time.sleep(interval)
try:
raw = tf.klines.ex_factors(chunk, **sdk_kwargs)
normalized = _normalize_adj_factor(raw)
if not normalized.is_empty():
all_dfs.append(normalized)
logger.debug("adj_factor chunk %d/%d: %d symbols", i + 1, len(chunks), len(chunk))
except Exception as e: # noqa: BLE001
logger.warning("adj_factor chunk %d failed: %s", i + 1, e)
if on_chunk_done:
on_chunk_done(i + 1, len(chunks))
if not all_dfs:
return 0, []
new_data = pl.concat(all_dfs, how="diagonal_relaxed") if len(all_dfs) > 1 else all_dfs[0]
# 提取受影响的 symbol 列表(合并前)
affected = new_data["symbol"].unique().to_list()
factor_dir = "adj_factor_etf" if asset_type == "etf" else "adj_factor"
out = repo.store.data_dir / factor_dir / "all.parquet"
out.parent.mkdir(parents=True, exist_ok=True)
if out.exists():
existing = pl.read_parquet(out)
before = existing.height
merged = pl.concat([existing, new_data]).unique(
subset=["symbol", "trade_date"], keep="last",
).sort(["symbol", "trade_date"])
merged.write_parquet(out)
added = merged.height - before
logger.info("adj_factor merged: %d total (+%d new), %d/%d symbols",
merged.height, added, new_data.height, len(symbols))
return added, affected
else:
new_data.sort(["symbol", "trade_date"]).write_parquet(out)
logger.info("adj_factor synced: %d rows (%d symbols)", new_data.height, len(symbols))
return new_data.height, affected
# ===== 分钟 K 同步 =====
CANONICAL_MINUTE_COLS = [
"symbol", "datetime", "open", "high", "low", "close", "volume", "amount",
]
def _normalize_minute(df_in, default_symbol: str | None = None) -> pl.DataFrame:
"""把 SDK 返回的分钟 K 数据规范成 canonical 列。"""
if df_in is None or len(df_in) == 0:
return pl.DataFrame()
if not isinstance(df_in, pl.DataFrame):
df = pl.from_pandas(df_in.reset_index() if hasattr(df_in, "reset_index") else df_in)
else:
df = df_in
rename_map = {
"ts_code": "symbol",
"vol": "volume",
"amt": "amount",
}
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
# datetime 列:优先用 timestamp(毫秒精度),其次 trade_time
if "timestamp" in df.columns:
df = df.with_columns(
pl.from_epoch("timestamp", time_unit="ms").alias("datetime"),
).drop("timestamp")
for drop_col in ("trade_time", "trade_date"):
if drop_col in df.columns:
df = df.drop(drop_col)
elif "trade_time" in df.columns:
df = df.rename({"trade_time": "datetime"})
if "trade_date" in df.columns:
df = df.drop("trade_date")
elif "trade_date" in df.columns:
df = df.rename({"trade_date": "datetime"})
if "symbol" not in df.columns and default_symbol is not None:
df = df.with_columns(pl.lit(default_symbol).alias("symbol"))
# 类型规范:统一转 Datetime('us')
if "datetime" in df.columns:
dt_type = df.schema["datetime"]
if not isinstance(dt_type, pl.Datetime) or dt_type.time_unit != "us":
df = df.with_columns(pl.col("datetime").cast(pl.Datetime("us"), strict=False))
for col in ("open", "high", "low", "close"):
if col in df.columns:
df = df.with_columns(pl.col(col).cast(pl.Float64, strict=False))
for col in ("volume", "amount"):
if col in df.columns:
df = df.with_columns(pl.col(col).cast(pl.Float64, strict=False))
keep = [c for c in CANONICAL_MINUTE_COLS if c in df.columns]
return df.select(keep)
def _datetime_to_ms(dt: datetime) -> int:
"""datetime → 毫秒时间戳 (供 SDK start_time / end_time 使用)。"""
return int(dt.timestamp() * 1000)
def sync_minute_batch(
symbols: list[str],
start_time: datetime | None = None,
end_time: datetime | None = None,
count: int | None = None,
batch_size: int | None = None,
rpm: int | None = None,
on_chunk_done: Callable[[int, int], None] | None = None,
) -> pl.DataFrame:
"""批量拉取多股分钟 K。
优先使用 start_time / end_time 区间, 确保所有标的覆盖同一时间段。
count 仅作为 fallback 保留。
on_chunk_done(current, total) 每个 chunk 完成后回调。
"""
tf = get_client()
out: list[pl.DataFrame] = []
interval = (60.0 / rpm) if rpm else 0
if batch_size is None:
chunks = [symbols]
else:
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
for i, chunk in enumerate(chunks):
if i > 0 and interval > 0 and len(chunks) > rpm:
time.sleep(interval)
try:
if start_time and end_time:
raw = tf.klines.batch(
chunk, period="1m",
start_time=_datetime_to_ms(start_time),
end_time=_datetime_to_ms(end_time),
count=10000,
as_dataframe=True, show_progress=False,
)
else:
raw = tf.klines.batch(chunk, period="1m", count=count or 1200,
as_dataframe=True, show_progress=False)
except Exception as e: # noqa: BLE001
logger.warning("minute batch fetch failed for %d symbols: %s", len(chunk), e)
continue
if isinstance(raw, dict):
for sym, sub in raw.items():
if sub is None or len(sub) == 0:
continue
out.append(_normalize_minute(sub, default_symbol=sym))
elif raw is not None and len(raw) > 0:
out.append(_normalize_minute(raw))
if on_chunk_done:
on_chunk_done(i + 1, len(chunks))
if not out:
return pl.DataFrame()
return pl.concat(out, how="diagonal_relaxed")
def fetch_minute_single(symbol: str, trade_date: date) -> pl.DataFrame:
"""从 TickFlow 实时拉取单股单日分钟 K(不写入本地)。"""
from datetime import datetime
start_time = datetime(trade_date.year, trade_date.month, trade_date.day, 9, 25, 0)
end_time = datetime(trade_date.year, trade_date.month, trade_date.day, 15, 5, 0)
tf = get_client()
try:
raw = tf.klines.batch(
[symbol], period="1m",
start_time=_datetime_to_ms(start_time),
end_time=_datetime_to_ms(end_time),
count=10000,
as_dataframe=True, show_progress=False,
)
except Exception as e:
logger.warning("fetch_minute_single(%s, %s) failed: %s", symbol, trade_date, e)
return pl.DataFrame()
if isinstance(raw, dict):
sub = raw.get(symbol)
return _normalize_minute(sub) if sub is not None and len(sub) > 0 else pl.DataFrame()
if raw is not None and len(raw) > 0:
return _normalize_minute(raw)
return pl.DataFrame()
def fetch_adj_factor_single(symbol: str) -> pl.DataFrame:
"""从 TickFlow 实时拉取单股除权因子(不写入本地), 用于单股 K 线即时前复权。
返回结构: symbol, trade_date, ex_factor (空 DataFrame 表示无除权事件或拉取失败)。
与 _apply_adj_factor / compute_enriched 的 factors 参数格式一致。
"""
tf = get_client()
try:
raw = tf.klines.ex_factors([symbol], as_dataframe=True, show_progress=False)
except Exception as e: # noqa: BLE001
logger.warning("fetch_adj_factor_single(%s) failed: %s", symbol, e)
return pl.DataFrame()
return _normalize_adj_factor(raw)
def _latest_minute_datetime(repo: KlineRepository) -> datetime | None:
"""本地分钟 K 数据的最新时间。"""
try:
res = repo.execute_one("SELECT max(datetime) FROM kline_minute")
if res and res[0]:
d = res[0]
if isinstance(d, datetime):
return d
return datetime.fromisoformat(str(d))
except Exception: # noqa: BLE001
pass
return None
def _cleanup_null_datetime_minute(repo: KlineRepository) -> None:
"""检测并清除 datetime 全为 null 的旧版分钟 K 数据(迁移用)。"""
minute_dir = repo.store.data_dir / "kline_minute"
if not minute_dir.exists():
return
try:
row = repo.execute_one(
"SELECT count(*) AS total, count(datetime) AS non_null FROM kline_minute"
)
if row and row[0] > 0 and (row[1] is None or row[1] == 0):
# 全部 datetime 为 null — 清除所有分钟 K parquet
n = 0
for f in minute_dir.rglob("*.parquet"):
f.unlink()
n += 1
logger.info("cleaned %d corrupted minute-K parquet files (null datetime)", n)
except Exception as e: # noqa: BLE001
logger.debug("minute cleanup check failed: %s", e)
def _migrate_symbol_to_date_partition(repo: KlineRepository) -> None:
"""将旧版 symbol= 分区迁移为 date= 分区。迁移完成后删除旧目录。"""
minute_dir = repo.store.data_dir / "kline_minute"
if not minute_dir.exists():
return
old_dirs = [d for d in minute_dir.iterdir() if d.is_dir() and d.name.startswith("symbol=")]
if not old_dirs:
return
logger.info("migrating %d symbol-partitioned minute-K dirs to date partition…", len(old_dirs))
all_frames: list[pl.DataFrame] = []
for sym_dir in old_dirs:
for pq in sym_dir.glob("*.parquet"):
try:
df = pl.read_parquet(pq)
if "datetime" in df.columns:
df = df.filter(pl.col("datetime").is_not_null())
if not df.is_empty():
all_frames.append(df)
except Exception: # noqa: BLE001
pass
if not all_frames:
# 数据全部不可用,直接删旧目录
for d in old_dirs:
d.mkdir(parents=True, exist_ok=True)
for f in d.rglob("*"):
if f.is_file():
f.unlink()
d.rmdir()
return
combined = pl.concat(all_frames, how="diagonal_relaxed")
combined = combined.unique(subset=["symbol", "datetime"], keep="last")
# 按日期写新分区
combined = combined.with_columns(pl.col("datetime").dt.date().alias("_trade_date"))
for day_df in combined.partition_by("_trade_date"):
trade_date = day_df["_trade_date"][0]
out = minute_dir / f"date={trade_date}" / "part.parquet"
out.parent.mkdir(parents=True, exist_ok=True)
day_df = day_df.drop("_trade_date").sort("symbol", "datetime")
day_df.write_parquet(out)
# 删旧目录
for d in old_dirs:
for f in d.rglob("*"):
if f.is_file():
f.unlink()
# 移除空目录
try:
d.rmdir()
except OSError:
pass
logger.info("minute-K migration done: %d rows migrated", combined.height)
def sync_and_persist_minute(
symbols: list[str],
repo: KlineRepository,
capset: CapabilitySet,
days: int = 5,
on_chunk_done: Callable[[int, int], None] | None = None,
) -> int:
"""同步分钟 K 并存到 Parquet(仅 raw,不前复权)。返回写入行数。
使用 start_time / end_time 区间拉取, 确保所有标的覆盖同一时间段。
on_chunk_done(current, total) 每个 chunk 完成后回调。
"""
if not symbols or not capset.has(Cap.KLINE_MINUTE_BATCH):
return 0
# 迁移:旧版 _normalize_minute 未转换 timestamp→datetime,导致全部 datetime 为 null
# 检测到后直接清除(这些数据无法使用)
_cleanup_null_datetime_minute(repo)
# 迁移:旧版按 symbol= 分区转为 date= 分区
_migrate_symbol_to_date_partition(repo)
now = datetime.now()
# 计算时间区间: 首次拉取回溯 N 天, 增量从最后数据时间开始
last_dt = _latest_minute_datetime(repo)
if last_dt:
start_time = last_dt
else:
start_time = now - timedelta(days=days)
end_time = now
lim = capset.limits(Cap.KLINE_MINUTE_BATCH)
batch_size = lim.batch if lim and lim.batch else 100
rpm = lim.rpm if lim else 30
df = sync_minute_batch(symbols, start_time=start_time, end_time=end_time,
batch_size=batch_size, rpm=rpm,
on_chunk_done=on_chunk_done)
if df.is_empty():
return 0
# 按日期分区写: data/kline_minute/date={YYYY-MM-DD}/part.parquet
df = df.with_columns(
pl.col("datetime").dt.date().alias("_trade_date")
)
written = 0
for day_df in df.partition_by("_trade_date"):
trade_date = day_df["_trade_date"][0]
out = repo.store.data_dir / "kline_minute" / f"date={trade_date}" / "part.parquet"
out.parent.mkdir(parents=True, exist_ok=True)
if out.exists():
existing = pl.read_parquet(out)
if "datetime" in existing.columns:
existing = existing.filter(pl.col("datetime").is_not_null())
day_df = pl.concat([existing, day_df.drop("_trade_date")]).unique(
subset=["symbol", "datetime"], keep="last",
)
else:
day_df = day_df.drop("_trade_date")
day_df = day_df.sort("symbol", "datetime")
day_df.write_parquet(out)
written += day_df.height
# 刷新视图
try:
d = repo.store.data_dir.as_posix()
repo.db.execute(
f"""CREATE OR REPLACE VIEW kline_minute AS
SELECT * FROM read_parquet('{d}/kline_minute/**/*.parquet', union_by_name=true)"""
)
except Exception as e: # noqa: BLE001
logger.warning("refresh kline_minute view failed: %s", e)
logger.info("minute K synced: %d rows (%d symbols)", written, len(symbols))
return written
@@ -0,0 +1,576 @@
"""市场总览数据装配(与 HTTP Request 解耦)。
本模块由 `app.api.overview._build_overview` 抽离而来,目的是让「大盘复盘」
等无 Request 的调用方(定时任务、复盘服务)也能复用同一套聚合逻辑。
行为与原 `_build_overview` 完全一致,仅把对 `request.app.state.{repo,
quote_service,depth_service}` 的依赖改为显式参数。
公共入口:
build_market_overview(repo, quote_service, depth_service, as_of)
"""
from __future__ import annotations
import math
import re
from datetime import date
from typing import Any
import polars as pl
from app.services.ext_data import ExtConfig, ExtConfigStore
from app.services.screener import ScreenerService
# ================================================================
# 常量(与 overview.py 保持同步;复盘复盘仅 A 股核心指数)
# ================================================================
CORE_INDEX_NAMES = {
"000001.SH": "上证指数",
"399001.SZ": "深证成指",
"399006.SZ": "创业板指",
"000680.SH": "科创综指",
}
CORE_INDEX_SYMBOLS = tuple(CORE_INDEX_NAMES.keys())
_DIMENSION_SEP = re.compile(r"[、,;|/\s]+")
# ================================================================
# 通用工具
# ================================================================
def _finite(v: Any) -> float | None:
if v is None:
return None
try:
f = float(v)
except (TypeError, ValueError):
return None
return f if math.isfinite(f) else None
def _json_safe(value: Any) -> Any:
if isinstance(value, dict):
return {k: _json_safe(v) for k, v in value.items()}
if isinstance(value, list):
return [_json_safe(v) for v in value]
if isinstance(value, float) and not math.isfinite(value):
return None
return value
def _board(symbol: str) -> str:
if symbol.endswith(".BJ"):
return "北交所"
if symbol.startswith(("300", "301")):
return "创业板"
if symbol.startswith(("688", "689")):
return "科创板"
if symbol.endswith(".SH"):
return "沪主板"
if symbol.endswith(".SZ"):
return "深主板"
return "其他"
def _score(value: float, low: float, high: float) -> int:
if high <= low:
return 50
return max(0, min(100, round((value - low) / (high - low) * 100)))
# ================================================================
# 指数行情(实时 quote_service 优先,回退 kline_index_daily SQL)
# ================================================================
def _quote_status(quote_service) -> dict:
qs = quote_service
if not qs:
return {"enabled": False, "running": False, "quote_age_ms": None, "is_trading_hours": False}
return qs.status()
def _index_quotes(repo, quote_service, as_of: date | None = None) -> list[dict]:
rows: list[dict] = []
if quote_service and as_of is None:
df = quote_service.get_index_quotes(list(CORE_INDEX_SYMBOLS))
if not df.is_empty():
rows = df.to_dicts()
if not rows and repo:
placeholders = ", ".join("?" for _ in CORE_INDEX_SYMBOLS)
try:
db_rows = repo.execute_all(
f"""
WITH ranked AS (
SELECT symbol, date, close,
row_number() OVER (PARTITION BY symbol ORDER BY date DESC) AS rn
FROM kline_index_daily
WHERE symbol IN ({placeholders})
AND (? IS NULL OR date <= ?)
), latest AS (
SELECT symbol,
max(CASE WHEN rn = 1 THEN date END) AS date,
max(CASE WHEN rn = 1 THEN close END) AS last_price,
max(CASE WHEN rn = 2 THEN close END) AS prev_close
FROM ranked
WHERE rn <= 2
GROUP BY symbol
)
SELECT symbol, date, last_price, prev_close
FROM latest
""",
[*CORE_INDEX_SYMBOLS, as_of, as_of],
)
except Exception: # noqa: BLE001
db_rows = []
for symbol, dt, last_price, prev_close in db_rows:
change_amount = None
change_pct = None
lp = _finite(last_price)
pc = _finite(prev_close)
if lp is not None and pc not in (None, 0):
change_amount = lp - pc
change_pct = change_amount / pc * 100
rows.append({
"symbol": symbol,
"name": CORE_INDEX_NAMES.get(symbol),
"date": str(dt) if dt else None,
"last_price": lp,
"close": lp,
"prev_close": pc,
"change_amount": change_amount,
"change_pct": change_pct,
})
by_symbol = {r.get("symbol"): r for r in rows}
out = []
for symbol in CORE_INDEX_SYMBOLS:
r = by_symbol.get(symbol, {"symbol": symbol})
out.append({
"symbol": symbol,
"name": r.get("name") or CORE_INDEX_NAMES[symbol],
"last_price": _finite(r.get("last_price") if r.get("last_price") is not None else r.get("close")),
"change_pct": _finite(r.get("change_pct")),
"change_amount": _finite(r.get("change_amount")),
})
return out
# ================================================================
# 扩展数据(行业 / 概念)维度聚合
# ================================================================
def _dimension_field(config: ExtConfig, kind: str) -> str | None:
candidates = ["概念", "concept", "theme"] if kind == "concept" else ["行业", "industry", "sector"]
for candidate in candidates:
needle = candidate.lower()
for field in config.fields:
haystack = f"{field.name} {field.label}".lower()
if needle in haystack:
return field.name
return None
def _ext_files(data_dir, config: ExtConfig) -> list[str]:
base = data_dir / "ext_data" / config.id
if config.mode == "timeseries":
root = base / "timeseries"
return [str(p) for p in sorted(root.rglob("*.parquet")) if p.is_file()]
return [str(p) for p in sorted(base.glob("*.parquet")) if p.is_file()]
def _read_ext_rows(data_dir, config: ExtConfig, dimension_field: str) -> list[dict]:
files = _ext_files(data_dir, config)
if not files:
return []
try:
df = pl.read_parquet(files, hive_partitioning=True)
except TypeError:
try:
df = pl.read_parquet(files)
except Exception: # noqa: BLE001
return []
except Exception: # noqa: BLE001
return []
if df.is_empty() or dimension_field not in df.columns:
return []
if config.mode == "timeseries" and "date" in df.columns:
latest = df.get_column("date").max()
if latest is not None:
df = df.filter(pl.col("date") == latest)
symbol_cols = ["symbol", "code", "股票代码", "代码"]
for mapping in (config.symbol_map, config.code_map):
if isinstance(mapping, dict) and mapping.get("type") == "mapped" and mapping.get("col"):
symbol_cols.append(str(mapping["col"]))
cols = []
for col in [dimension_field, *symbol_cols]:
if col in df.columns and col not in cols:
cols.append(col)
return df.select(cols).to_dicts()
def _dimension_values(raw: Any) -> list[str]:
if raw is None:
return []
values = [v.strip() for v in _DIMENSION_SEP.split(str(raw).strip()) if v.strip()]
return values
def _symbol_keys(row: dict, config: ExtConfig) -> list[str]:
fields = ["symbol", "code", "股票代码", "代码"]
for mapping in (config.symbol_map, config.code_map):
if isinstance(mapping, dict) and mapping.get("type") == "mapped" and mapping.get("col"):
fields.append(str(mapping["col"]))
keys: list[str] = []
for field in fields:
raw = row.get(field)
if raw is None:
continue
text = str(raw).strip().upper()
if not text:
continue
keys.append(text)
if "." in text:
keys.append(text.split(".", 1)[0])
return keys
def _dimension_rank(rows: list[dict], repo, kind: str, limit: int = 5, level: int | None = None) -> dict:
if not rows:
return {"leading": [], "lagging": []}
quote_map: dict[str, dict] = {}
for row in rows:
symbol = str(row.get("symbol") or "").strip().upper()
if not symbol:
continue
quote_map[symbol] = row
quote_map[symbol.split(".", 1)[0]] = row
store = ExtConfigStore(repo.store.data_dir)
groups: dict[str, dict[str, dict]] = {}
for config in store.load_all():
field = _dimension_field(config, kind)
if not field:
continue
for ext_row in _read_ext_rows(repo.store.data_dir, config, field):
quote = None
for key in _symbol_keys(ext_row, config):
quote = quote_map.get(key)
if quote:
break
if not quote:
continue
symbol = str(quote.get("symbol") or "")
for value in _dimension_values(ext_row.get(field)):
# 行业按 "-" 拆分级: "银行-银行-股份制银行" → level=2 取"银行"(二级)
if level is not None and "-" in value:
parts = value.split("-")
value = parts[level - 1] if level <= len(parts) else parts[-1]
groups.setdefault(value, {})[symbol] = quote
items = []
for name, by_symbol in groups.items():
stocks = list(by_symbol.values())
changes = [_finite(s.get("change_pct")) for s in stocks]
changes = [v for v in changes if v is not None]
if not changes:
continue
leader = max(stocks, key=lambda s: _finite(s.get("change_pct")) or -999)
items.append({
"name": name,
"count": len(stocks),
"avg_pct": sum(changes) / len(changes),
"up_count": sum(1 for v in changes if v > 0),
"down_count": sum(1 for v in changes if v < 0),
"amount": sum(_finite(s.get("amount")) or 0 for s in stocks),
"leader": {
"symbol": leader.get("symbol"),
"name": leader.get("name"),
"change_pct": _finite(leader.get("change_pct")),
},
})
leading = sorted(items, key=lambda x: x["avg_pct"], reverse=True)[:limit]
lagging = sorted(items, key=lambda x: x["avg_pct"])[:limit]
return {"leading": leading, "lagging": lagging}
# ================================================================
# Top 行 / 涨跌幅分桶
# ================================================================
def _top_rows(rows: list[dict], key: str, descending: bool, limit: int = 8) -> list[dict]:
filtered = [r for r in rows if _finite(r.get(key)) is not None]
filtered.sort(key=lambda r: _finite(r.get(key)) or 0, reverse=descending)
return [
{
"symbol": r.get("symbol"),
"name": r.get("name"),
"close": _finite(r.get("close")),
"change_pct": _finite(r.get("change_pct")),
"amount": _finite(r.get("amount")),
"turnover_rate": _finite(r.get("turnover_rate")),
"board": _board(str(r.get("symbol") or "")),
}
for r in filtered[:limit]
]
def _pct_band_rows(values: list[float]) -> list[dict]:
bands = [
("<-5%", None, -0.05),
("-5~-3%", -0.05, -0.03),
("-3~-1%", -0.03, -0.01),
("-1~0%", -0.01, 0),
("0~1%", 0, 0.01),
("1~3%", 0.01, 0.03),
("3~5%", 0.03, 0.05),
(">5%", 0.05, None),
]
total = len(values) or 1
out = []
for label, low, high in bands:
count = 0
for v in values:
if low is None and v < high:
count += 1
elif high is None and v >= low:
count += 1
elif low is not None and high is not None and low <= v < high:
count += 1
out.append({"label": label, "count": count, "pct": count / total * 100})
return out
# ================================================================
# 主装配入口
# ================================================================
def build_market_overview(
repo,
quote_service=None,
depth_service=None,
as_of: date | None = None,
) -> dict:
"""装配市场总览(与原 overview._build_overview 行为一致)。
Args:
repo: KlineRepository(必填)。
quote_service: QuoteService(可选;实时指数行情来源)。
depth_service: DepthService(可选;五档封板修正)。
as_of: 指定日期,None 则取最新有数据日。
"""
svc = ScreenerService(repo)
as_of = as_of or svc.latest_date()
status = _quote_status(quote_service)
indices = _index_quotes(repo, quote_service, as_of)
if not as_of:
return {
"as_of": None,
"quote_status": status,
"indices": indices,
"breadth": {"total": 0, "up": 0, "down": 0, "flat": 0, "up_pct": 0, "down_pct": 0},
"amount": {"total": 0, "avg": 0},
"boards": [],
"limit": {"limit_up": 0, "broken": 0, "failed": 0, "limit_down": 0, "max_boards": 0, "tiers": []},
"distribution": [],
"trend": {"above_ma5": 0, "above_ma20": 0, "above_ma60": 0, "above_ma5_pct": 0, "above_ma20_pct": 0, "above_ma60_pct": 0, "new_high": 0, "new_low": 0},
"activity": {"avg_turnover": 0, "high_turnover": 0, "high_vol_ratio": 0, "vol_ratio": 1},
"radar": [],
"emotion": {"score": 50, "label": "暂无"},
"top_gainers": [],
"top_losers": [],
"turnover_leaders": [],
"active_leaders": [],
"concept_rank": {"leading": [], "lagging": []},
"industry_rank": {"leading": [], "lagging": []},
}
df = svc._load_enriched_for_date(as_of)
if df.is_empty():
rows: list[dict] = []
else:
cols = [
"symbol", "name", "close", "change_pct", "amount", "turnover_rate", "volume",
"vol_ratio_5d", "consecutive_limit_ups", "signal_limit_up", "signal_broken_limit_up", "signal_limit_down",
"ma5", "ma20", "ma60", "high_60d", "low_60d", "signal_n_day_high", "signal_n_day_low",
]
df = df.select([c for c in cols if c in df.columns])
rows = df.to_dicts()
# 过滤真停牌(volume=0 且 change_pct=0),保留有涨跌幅的浮点误差股以对齐同花顺口径
if rows and "volume" in rows[0]:
rows = [r for r in rows
if (_finite(r.get("volume")) or 0) > 0
or (_finite(r.get("change_pct")) or 0) != 0]
total = len(rows)
up = sum(1 for r in rows if (_finite(r.get("change_pct")) or 0) > 0)
down = sum(1 for r in rows if (_finite(r.get("change_pct")) or 0) < 0)
flat = max(0, total - up - down)
up_pct = up / total * 100 if total else 0
down_pct = down / total * 100 if total else 0
amounts = [_finite(r.get("amount")) or 0 for r in rows]
total_amount = sum(amounts)
avg_amount = total_amount / total if total else 0
pct_values = [_finite(r.get("change_pct")) for r in rows]
pct_values = [v for v in pct_values if v is not None]
avg_pct = sum(pct_values) / len(pct_values) if pct_values else 0
median_pct = sorted(pct_values)[len(pct_values) // 2] if pct_values else 0
strong_up = sum(1 for v in pct_values if v >= 0.03)
strong_down = sum(1 for v in pct_values if v <= -0.03)
limit_up = sum(1 for r in rows if bool(r.get("signal_limit_up")) or (_finite(r.get("consecutive_limit_ups")) or 0) > 0)
broken = sum(1 for r in rows if bool(r.get("signal_broken_limit_up")))
limit_down = sum(1 for r in rows if bool(r.get("signal_limit_down")))
max_boards = max([int(_finite(r.get("consecutive_limit_ups")) or 0) for r in rows], default=0)
# 五档 sealed 修正: 假涨停/假跌停不计入(需 Pro+ depth5.batch 能力)
sealed_ready = False
fake_up = 0
fake_down = 0
if depth_service:
up_map = depth_service.get_sealed_map(as_of, is_down=False)
down_map = depth_service.get_sealed_map(as_of, is_down=True)
sealed_ready = bool(up_map or down_map) and depth_service.is_sealed_ready(as_of)
if up_map:
fake_up = sum(1 for v in up_map.values() if v.get("sealed") is False)
if down_map:
fake_down = sum(1 for v in down_map.values() if v.get("sealed") is False)
if sealed_ready:
limit_up = max(0, limit_up - fake_up)
limit_down = max(0, limit_down - fake_down)
seal_rate = limit_up / (limit_up + broken) * 100 if (limit_up + broken) > 0 else 0
def above_ma_count(ma_key: str) -> int:
return sum(1 for r in rows if (_finite(r.get("close")) is not None and _finite(r.get(ma_key)) is not None and (_finite(r.get("close")) or 0) >= (_finite(r.get(ma_key)) or 0)))
above_ma5 = above_ma_count("ma5")
above_ma20 = above_ma_count("ma20")
above_ma60 = above_ma_count("ma60")
new_high = sum(1 for r in rows if bool(r.get("signal_n_day_high")) or (_finite(r.get("close")) is not None and _finite(r.get("high_60d")) is not None and (_finite(r.get("close")) or 0) >= (_finite(r.get("high_60d")) or 0)))
new_low = sum(1 for r in rows if bool(r.get("signal_n_day_low")) or (_finite(r.get("close")) is not None and _finite(r.get("low_60d")) is not None and (_finite(r.get("close")) or 0) <= (_finite(r.get("low_60d")) or 0)))
turnovers = [_finite(r.get("turnover_rate")) for r in rows]
turnovers = [v for v in turnovers if v is not None]
avg_turnover = sum(turnovers) / len(turnovers) if turnovers else 0
high_turnover = sum(1 for v in turnovers if v >= 5)
boards_map: dict[str, dict] = {}
for r in rows:
b = _board(str(r.get("symbol") or ""))
item = boards_map.setdefault(b, {"board": b, "count": 0, "up": 0, "down": 0, "amount": 0.0})
item["count"] += 1
change = _finite(r.get("change_pct")) or 0
if change > 0:
item["up"] += 1
elif change < 0:
item["down"] += 1
item["amount"] += _finite(r.get("amount")) or 0
boards = sorted(boards_map.values(), key=lambda x: x["amount"], reverse=True)
for b in boards:
count = b["count"] or 1
b["up_pct"] = b["up"] / count * 100
tiers_map: dict[int, int] = {}
for r in rows:
n = int(_finite(r.get("consecutive_limit_ups")) or 0)
if n > 0:
tiers_map[n] = tiers_map.get(n, 0) + 1
tiers = [{"boards": k, "count": v} for k, v in sorted(tiers_map.items(), key=lambda item: -item[0])]
index_changes = [_finite(r.get("change_pct")) for r in indices]
index_changes = [v for v in index_changes if v is not None]
avg_index_pct = sum(index_changes) / len(index_changes) if index_changes else 0
vol_ratios = [_finite(r.get("vol_ratio_5d")) for r in rows]
vol_ratios = [v for v in vol_ratios if v is not None]
avg_vol_ratio = sum(vol_ratios) / len(vol_ratios) if vol_ratios else 1
high_vol_ratio = sum(1 for v in vol_ratios if v >= 1.5)
concept_rank = _dimension_rank(rows, repo, "concept")
industry_rank = _dimension_rank(rows, repo, "industry", level=2)
strong_diff_pct = (strong_up - strong_down) / total * 100 if total else 0
high_vol_pct = high_vol_ratio / total * 100 if total else 0
strong_down_pct = strong_down / total * 100 if total else 0
tier2_count = sum(t["count"] for t in tiers if t["boards"] >= 2)
mainline_items = [*concept_rank["leading"][:3], *industry_rank["leading"][:3]]
mainline_avg = max([_finite(item.get("avg_pct")) or 0 for item in mainline_items], default=0)
mainline_cover_pct = max([(_finite(item.get("count")) or 0) / total * 100 for item in mainline_items], default=0) if total else 0
mainline_score = round(_score(mainline_avg, -0.005, 0.03) * 0.65 + _score(mainline_cover_pct, 1, 12) * 0.35) if mainline_items else 50
radar = [
{"key": "index", "label": "指数", "value": _score(avg_index_pct, -2.5, 2.5)},
{"key": "profit", "label": "赚钱", "value": round(_score(up_pct, 20, 80) * 0.45 + _score(avg_pct, -0.02, 0.02) * 0.25 + _score(median_pct, -0.02, 0.02) * 0.20 + _score(strong_diff_pct, -8, 8) * 0.10)},
{"key": "money", "label": "量能", "value": round(_score(avg_vol_ratio, 0.6, 1.8) * 0.70 + _score(high_vol_pct, 2, 12) * 0.30)},
{"key": "speculation", "label": "投机", "value": round(_score(limit_up, 5, 90) * 0.25 + _score(seal_rate, 30, 85) * 0.35 + _score(max_boards, 1, 8) * 0.25 + _score(tier2_count, 0, 30) * 0.15)},
{"key": "resilience", "label": "抗跌", "value": 100 - round(_score(down_pct, 20, 80) * 0.55 + _score(strong_down_pct, 1, 12) * 0.45)},
{"key": "mainline", "label": "主线", "value": mainline_score},
]
emotion_score = round(sum(r["value"] for r in radar) / len(radar)) if radar else 50
if emotion_score >= 70:
emotion_label = "强势"
elif emotion_score >= 55:
emotion_label = "偏暖"
elif emotion_score >= 45:
emotion_label = "震荡"
elif emotion_score >= 30:
emotion_label = "偏冷"
else:
emotion_label = "冰点"
return _json_safe({
"as_of": str(as_of),
"quote_status": status,
"indices": indices,
"breadth": {
"total": total,
"up": up,
"down": down,
"flat": flat,
"up_pct": up_pct,
"down_pct": down_pct,
"avg_pct": avg_pct,
"median_pct": median_pct,
"strong_up": strong_up,
"strong_down": strong_down,
},
"amount": {"total": total_amount, "avg": avg_amount},
"boards": boards,
"limit": {"limit_up": limit_up, "broken": broken, "failed": 0, "limit_down": limit_down, "max_boards": max_boards, "seal_rate": seal_rate, "tiers": tiers, "sealed_ready": sealed_ready, "fake_up": fake_up, "fake_down": fake_down},
"distribution": _pct_band_rows(pct_values),
"trend": {
"above_ma5": above_ma5,
"above_ma20": above_ma20,
"above_ma60": above_ma60,
"above_ma5_pct": above_ma5 / total * 100 if total else 0,
"above_ma20_pct": above_ma20 / total * 100 if total else 0,
"above_ma60_pct": above_ma60 / total * 100 if total else 0,
"new_high": new_high,
"new_low": new_low,
},
"activity": {
"avg_turnover": avg_turnover,
"high_turnover": high_turnover,
"high_vol_ratio": high_vol_pct,
"vol_ratio": avg_vol_ratio,
},
"radar": radar,
"emotion": {"score": emotion_score, "label": emotion_label},
"top_gainers": _top_rows(rows, "change_pct", True),
"top_losers": _top_rows(rows, "change_pct", False),
"turnover_leaders": _top_rows(rows, "amount", True),
"active_leaders": _top_rows(rows, "turnover_rate", True),
"concept_rank": concept_rank,
"industry_rank": industry_rank,
})
+343
View File
@@ -0,0 +1,343 @@
"""AI 大盘复盘 —— 流式 LLM 复盘生成。
复刻 stock_analyzer.py 的 NDJSON 流式协议(meta/delta/error/done),
将「市场总览」聚合数据交给 LLM 生成结构化复盘报告。
数据来源:services.market_overview_builder.build_market_overview
(与 GET /api/overview/market 同源,保证复盘与看板数据口径一致)。
流式协议(与 stock_analyzer / financial_analyzer 一致,前端解析无差异):
{"type":"meta", "as_of", "emotion_score", "emotion_label", "summary"}
{"type":"delta","content":"..."} 逐 chunk 文本
{"type":"error","message":"..."}
{"type":"done"}
"""
from __future__ import annotations
import json
import logging
from datetime import date
from typing import AsyncIterator
from app.services.market_overview_builder import build_market_overview
logger = logging.getLogger(__name__)
# 指数简称映射:摘要里用简称(上/深/创/科),全称太长列表放不下。与前端 INDEX_SHORT 对齐。
_INDEX_SHORT = {
"上证指数": "",
"深证成指": "",
"创业板指": "",
"科创综指": "",
"科创50": "",
}
# ================================================================
# 系统提示词(市场策略师人格 + 固定七节模板)
# ================================================================
_SYSTEM_PROMPT = """你是一位拥有 15 年 A 股一线实战经验的资深市场策略师,擅长从指数结构、涨跌家数、连板梯队、板块轮动与资金情绪中提炼交易主线,产出可直接指导次日仓位与节奏的盘后复盘报告。
## 输出规范
用 **Markdown** 格式输出,严格遵循以下结构。不要输出任何 JSON 或代码块,直接输出 Markdown 正文。
### 1. 🎯 一句话定调(1-2 句)
用一句话概括今日市场的**核心矛盾与状态**(如"放量普涨、情绪修复,主线围绕科技扩散"/"指数虚高、个股杀跌,赚钱效应冰点")。结尾用【明日基调:进攻 / 均衡 / 防守】给出明确倾向。
### 2. 📊 盘面总览
- 三大指数(上证/深证/创业板)表现:谁强谁弱、量能配合
- 涨跌家数、涨停/跌停/炸板结构、两市成交额(放量/缩量判断)
- 情绪温度(强势/偏暖/震荡/偏冷/冰点)及一句话依据
### 3. 📈 指数结构
谁在护盘、谁在拖累;指数是否同步;关键支撑/压力位(基于当日点位推断);是否存在量价背离。
### 4. 🔥 板块主线
- 领涨板块:背后的逻辑(消息/业绩/资金/技术)、持续性判断、是否形成可交易主线
- 领跌板块:风险信号、是否扩散
- 连板梯队与投机情绪:最高连板、封板率、炸板率反映的资金激进程度
### 5. 💰 资金与情绪
成交额结构(增量/存量)、市场宽度(上涨占比、站上均线占比)、量能指标(量比)解读;风险偏好是修复还是转弱。
### 6. 📰 消息催化
结合提供的近期新闻,提炼真正影响明日交易节奏的催化或扰动,明确区分"已兑现""待发酵"。**若无新闻数据,则直接从量价异动推断可能的催化逻辑并给出结论,不要标注"[推断]"之类的过程标签,更不要编造具体消息。**
### 7. 🎯 明日交易计划
- 进攻 / 均衡 / 防守:基于今日盘面给出次日基调
- 仓位区间建议(轻仓/半仓/重仓的粗略指引)
- 关注方向(领涨延续 / 低吸 / 反包)与回避方向(高位滞涨 / 杀跌扩散)
- 一个明确的触发失效条件(如"若上证跌破 X 点则转为防守")
### 8. ⚠️ 风险提示
列出需要重点盯的风险点(如量能跟不上、外资流出、连板断层等)。末尾附一行:
"> ⚠️ 本报告由 AI 基于公开行情数据生成,仅供参考,不构成任何投资建议。交易有风险,入市需谨慎。"
## 分析准则(务必遵守)
0. **只输出结论,不输出思考过程**:禁止复述你的分析步骤或方法论。不要写"我先按...做结构化复盘""接下来看...""基于上述数据我认为"这类元话语——直接给结论。读者要的是复盘结果,不是你怎么推导出来的。
1. **数据说话**:每个判断引用具体数值,严禁空泛套话("情绪回暖"必须改成"涨停 68 家较前日 +22,封板率 75%")
2. **诚实中立**:看多就写多,看空就写空,不要骑墙;数据不支持时直言无法判断
3. **结构优先**:先看指数同步性与量能结构,再看板块与情绪,最后才是消息
4. **不重复数字**:正文负责解读表格数据背后的含义,不要照抄罗列已提供的大段原始数字
5. **风险前置**:任何进攻建议都要配触发失效条件
6. **简明实战**:用交易员能扫读的密度输出,总字数 1200-2000 字,重在可执行
现在请基于下方数据进行复盘。"""
# ================================================================
# 用户消息构建(精简切片,控制 token)
# ================================================================
def _fmt_pct(v, suffix="%") -> str:
if v is None:
return ""
return f"{v:+.2f}{suffix}" if suffix else f"{v:.2f}"
def _build_indices_block(overview: dict) -> str:
"""指数行情精简块。"""
indices = overview.get("indices") or []
if not indices:
return "(暂无指数)"
lines = []
for idx in indices:
name = idx.get("name") or idx.get("symbol")
price = idx.get("last_price")
chg = idx.get("change_pct")
price_s = f"{price:.2f}" if price is not None else ""
lines.append(f"- {name}: {price_s} {_fmt_pct(chg)}")
return "\n".join(lines)
def _build_breadth_block(overview: dict) -> str:
b = overview.get("breadth") or {}
amt = overview.get("amount") or {}
lim = overview.get("limit") or {}
tr = overview.get("trend") or {}
act = overview.get("activity") or {}
total_amount = amt.get("total") or 0
# 成交额单位换算为亿元(原始为元)
amount_yi = total_amount / 1e8 if total_amount else 0
lines = [
f"- 上涨/下跌/平盘: {b.get('up',0)} / {b.get('down',0)} / {b.get('flat',0)}"
f" (上涨占比 {b.get('up_pct',0):.1f}%)",
f"- 涨停/炸板/跌停: {lim.get('limit_up',0)} / {lim.get('broken',0)} / {lim.get('limit_down',0)}"
f" (封板率 {lim.get('seal_rate',0):.0f}%, 最高连板 {lim.get('max_boards',0)})",
]
if lim.get("tiers"):
tiers_str = "".join(f"{t['boards']}板×{t['count']}" for t in lim["tiers"][:5])
lines.append(f"- 连板梯队: {tiers_str}")
lines.append(f"- 两市成交额: {amount_yi:.0f} 亿元")
lines.append(
f"- 均线站位: MA5 {tr.get('above_ma5_pct',0):.0f}% / "
f"MA20 {tr.get('above_ma20_pct',0):.0f}% / MA60 {tr.get('above_ma60_pct',0):.0f}%"
)
lines.append(
f"- 量能: 平均换手 {act.get('avg_turnover',0):.2f}%, "
f"量比5日均 {act.get('vol_ratio',1):.2f}"
)
return "\n".join(lines)
def _build_sector_block(rank: dict, label: str) -> str:
"""板块排名精简块(领涨/领跌 top5)。"""
if not rank:
return f"### {label}\n(暂无数据)"
def _fmt(items):
if not items:
return ""
return "".join(
f"{it.get('name')}({(it.get('avg_pct') or 0)*100:+.2f}%,领涨:{it.get('leader',{}).get('name','')})"
for it in items[:5]
)
return (
f"- 领涨{label}: {_fmt(rank.get('leading'))}\n"
f"- 领跌{label}: {_fmt(rank.get('lagging'))}"
)
def _build_emotion_block(overview: dict) -> str:
emo = overview.get("emotion") or {}
radar = overview.get("radar") or []
score = emo.get("score", 50)
label = emo.get("label", "")
lines = [f"- 情绪温度: {score} ({label})"]
if radar:
dims = "".join(f"{r.get('label')}{r.get('value',0)}" for r in radar)
lines.append(f"- 六维雷达: {dims}")
return "\n".join(lines)
def _build_user_prompt(overview: dict, news: list[dict], focus: str) -> str:
"""构建用户消息:复盘日期 + 市场数据精简切片 + 新闻 + 关注点。"""
as_of = overview.get("as_of") or "今日"
parts: list[str] = [
f"复盘日期: {as_of}",
"",
"## 主要指数",
_build_indices_block(overview),
"",
"## 盘面数据",
_build_breadth_block(overview),
"",
"## 市场情绪",
_build_emotion_block(overview),
"",
"## 概念板块排名",
_build_sector_block(overview.get("concept_rank"), "概念"),
"",
"## 行业板块排名",
_build_sector_block(overview.get("industry_rank"), "行业"),
]
if news:
news_lines = []
for i, n in enumerate(news[:8], 1):
title = (n.get("title") or "").strip()
snippet = (n.get("snippet") or "").strip()
source = (n.get("source") or "").strip()
pub = (n.get("published_date") or "").strip()
meta = " / ".join(p for p in (source, pub) if p)
news_lines.append(f"{i}. {title} ({meta})\n {snippet}" if meta else f"{i}. {title}\n {snippet}")
parts.extend(["", "## 近期市场新闻", "\n".join(news_lines)])
else:
parts.extend([
"",
"## 近期市场新闻",
"(暂无新闻数据:本功能新闻检索能力将在后续版本接入。"
"消息催化一节请直接从量价异动给出可能的催化逻辑结论,不要编造具体消息,也不要复述本说明。)",
])
if focus.strip():
parts.extend(["", f"本次复盘请特别关注: {focus.strip()}"])
return "\n".join(parts)
# ================================================================
# 摘要生成(供 meta 事件 / 历史报告 summary)
# ================================================================
def _recap_summary(overview: dict) -> str:
"""一句话摘要(供 meta 事件与历史列表展示)。
指数用简称(上/深/创/科),与前端摘要条一致,避免列表里全称放不下。
"""
indices = overview.get("indices") or []
emo = overview.get("emotion") or {}
lim = overview.get("limit") or {}
amt = overview.get("amount") or {}
total_amount = (amt.get("total") or 0) / 1e8
idx_str = "".join(
f"{_INDEX_SHORT.get(i.get('name') or '', i.get('name') or '')}{(i.get('change_pct') or 0):+.2f}%"
for i in indices[:4]
) or "指数缺失"
return (
f"{idx_str} | 情绪{emo.get('score',50)}({emo.get('label','')}) | "
f"涨停{lim.get('limit_up',0)} | 成交{total_amount:.0f}亿"
)
# ================================================================
# 流式主入口
# ================================================================
async def recap_market_stream(
repo,
quote_service=None,
depth_service=None,
as_of: date | None = None,
focus: str = "",
news: list[dict] | None = None,
) -> AsyncIterator[str]:
"""流式大盘复盘:yield 出每个 NDJSON 事件。
Args:
repo: KlineRepository(必填)。
quote_service / depth_service: 可选,数据装配依赖。
as_of: 复盘日期,None 取最新有数据日。
focus: 用户追加的复盘关注点。
news: 预检索的新闻列表(P1 不传,留 None 走降级说明;P3 由 news_search 注入)。
"""
# 1. 装配市场总览
overview = build_market_overview(repo, quote_service, depth_service, as_of)
as_of_str = overview.get("as_of")
if not as_of_str:
yield json.dumps({
"type": "error",
"message": "暂无市场数据,请先在「数据」页同步日 K 与指数后再复盘",
}, ensure_ascii=False)
return
emo = overview.get("emotion") or {}
# 2. meta 事件(前端据此先渲染信号灯/看板)
yield json.dumps({
"type": "meta",
"as_of": as_of_str,
"emotion_score": emo.get("score", 50),
"emotion_label": emo.get("label", ""),
"summary": _recap_summary(overview),
}, ensure_ascii=False)
# 3+4. 构建 prompt + 流式调用 LLM(整体 try-except,任何异常 yield error,避免前端卡死)
try:
from app.services.ai_provider import stream_ai_text
user_prompt = _build_user_prompt(overview, news or [], focus)
async for delta in stream_ai_text(
[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
],
temperature=0.5,
max_tokens=4500,
):
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
except Exception as e: # noqa: BLE001
logger.exception("AI market recap failed for %s: %s", as_of_str, e)
yield json.dumps({"type": "error", "message": f"AI 复盘失败: {e}"}, ensure_ascii=False)
return
yield json.dumps({"type": "done"}, ensure_ascii=False)
async def recap_market_once(
repo,
quote_service=None,
depth_service=None,
as_of: date | None = None,
focus: str = "",
news: list[dict] | None = None,
) -> tuple[str | None, dict]:
"""非流式版本(供定时任务调用):累积全部 delta,返回 (content, meta)。
content 为完整 Markdown 文本;失败时为 None。
meta 含 as_of / emotion_score / emotion_label / summary(即使失败也尽量回填)。
"""
content_parts: list[str] = []
meta: dict = {"as_of": as_of.isoformat() if as_of else None}
async for evt in recap_market_stream(repo, quote_service, depth_service, as_of, focus, news):
try:
obj = json.loads(evt)
except Exception: # noqa: BLE001
continue
t = obj.get("type")
if t == "meta":
meta = obj
elif t == "delta":
content_parts.append(obj.get("content", ""))
elif t == "error":
logger.warning("market recap error event: %s", obj.get("message"))
return None, meta
return "".join(content_parts), meta
@@ -0,0 +1,92 @@
"""AI 大盘复盘报告持久化存储。
与 stock_reports.py(个股分析报告)/ ai_reports.py(财务分析报告)完全独立 ——
单独的文件、字段、上限,互不影响。刻意不复用,避免引入 kind 判别字段与分支
(解耦 > 抽象)。
存储位置: data/user_data/ai_market_recaps.json (数组,按 created_at 降序)
保留最近 MAX_REPORTS 条;超出自动裁剪最旧的。
每条报告结构:
{
"id": "mkr_xxx", # 唯一 id(market-recap-report)
"as_of": "2026-06-27", # 复盘日期
"focus": "", # 用户追加的关心点(可为空)
"content": "# ...markdown", # 报告正文
"summary": "三大指数齐涨...", # 一句话摘要
"emotion_score": 68, # 情绪分(0-100, 复盘生成时的市场情绪雷达均分)
"emotion_label": "偏暖", # 情绪标签(强势/偏暖/震荡/偏冷/冰点)
"created_at": "2026-06-27T15:35:00"
}
"""
from __future__ import annotations
import json
import logging
import time
from pathlib import Path
logger = logging.getLogger(__name__)
MAX_REPORTS = 20
def _path() -> Path:
from app.config import settings
p = settings.data_dir / "user_data" / "ai_market_recaps.json"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def list_reports() -> list[dict]:
"""返回全部报告(按 created_at 降序)。"""
p = _path()
if not p.exists():
return []
try:
data = json.loads(p.read_text(encoding="utf-8"))
if isinstance(data, list):
return sorted(data, key=lambda r: r.get("created_at", ""), reverse=True)
except Exception as e: # noqa: BLE001
logger.warning("ai_market_recaps.json malformed: %s", e)
return []
def _save_all(reports: list[dict]) -> None:
"""全量写入(裁剪到 MAX_REPORTS)。"""
reports.sort(key=lambda r: r.get("created_at", ""), reverse=True)
if len(reports) > MAX_REPORTS:
reports = reports[:MAX_REPORTS]
_path().write_text(
json.dumps(reports, indent=2, ensure_ascii=False), encoding="utf-8",
)
def save_report(report: dict) -> dict:
"""新增一条报告并持久化。返回保存后的报告(含 id / created_at)。"""
reports = list_reports()
if not report.get("id"):
report["id"] = f"mkr_{int(time.time() * 1000)}"
if not report.get("created_at"):
report["created_at"] = _now_iso()
reports.append(report)
_save_all(reports)
logger.info("Market recap saved: %s (as_of=%s), total %d",
report.get("id"), report.get("as_of"), len(reports))
return report
def delete_report(report_id: str) -> bool:
"""删除指定报告。返回是否删除成功。"""
reports = list_reports()
before = len(reports)
reports = [r for r in reports if r.get("id") != report_id]
if len(reports) < before:
_save_all(reports)
return True
return False
def _now_iso() -> str:
from datetime import datetime
return datetime.now().isoformat(timespec="seconds")
@@ -0,0 +1,136 @@
"""系统通知适配器 — 三平台原生通知中心。
职责: 把后端产生的告警事件推送到操作系统通知中心。
窗口最小化 / 被遮挡 / 后台运行时都能弹通知 (不依赖前端 WebView)。
平台实现:
- Windows: winotify (进现代操作中心, 支持图标)
- macOS: osascript (系统已内置, 无需额外依赖)
- Linux: notify-send (系统已内置) / plyer 兜底
设计: 失败静默降级, 绝不因通知失败阻断告警主流程 (落盘 / SSE 推送)。
通知去重不在本层做, 复用 MonitorRuleEngine 的 cooldown 逻辑。
"""
from __future__ import annotations
import logging
import subprocess
import sys
from pathlib import Path
logger = logging.getLogger(__name__)
# 单次通知最长字符 (避免某些平台截断报错)
_MAX_LEN = 200
# 避免重复探测平台能力, 缓存一次
_backend_cache: str | None = None
def _detect_backend() -> str | None:
"""探测当前平台可用的通知后端。"""
global _backend_cache
if _backend_cache is not None:
return _backend_cache if _backend_cache != "none" else None
backend = None
if sys.platform == "win32":
try:
import winotify # type: ignore[import-not-found] # noqa: F401
backend = "winotify"
except ImportError:
logger.debug("winotify 不可用, Windows 通知降级")
backend = None
elif sys.platform == "darwin":
backend = "osascript"
elif sys.platform.startswith("linux"):
backend = "notify-send"
_backend_cache = backend or "none"
return backend
def _truncate(text: str) -> str:
"""截断超长文本, 避免平台通知上限报错。"""
text = (text or "").strip()
return text[:_MAX_LEN] + ("" if len(text) > _MAX_LEN else "")
def notify(title: str, message: str, icon: Path | None = None) -> bool:
"""推送一条系统通知。
Args:
title: 通知标题
message: 通知正文
icon: 可选图标路径 (部分平台支持)
Returns:
True=成功送达, False=失败或无可用后端。
失败静默, 不抛异常 (通知是辅助通道, 不能阻断告警主流程)。
"""
title = _truncate(title)
message = _truncate(message)
if not title:
return False
backend = _detect_backend()
if backend is None:
return False
try:
if backend == "winotify":
return _notify_winotify(title, message)
if backend == "osascript":
return _notify_osascript(title, message)
if backend == "notify-send":
return _notify_notify_send(title, message, icon)
except Exception as e: # noqa: BLE001
logger.debug("系统通知失败 (%s): %s", backend, e)
return False
return False
def _notify_winotify(title: str, message: str) -> bool:
"""Windows 通知 (winotify) — 进现代操作中心。"""
from winotify import Notifier # type: ignore[import-not-found]
Notifier().create_notification(
title=title,
msg=message,
# winotify 要求 duration 为 "short" 或 "long"
duration="short",
# 无可点击动作 (桌面版不实现"点击回到窗口"的复杂交互)
).show()
return True
def _notify_osascript(title: str, message: str) -> bool:
"""macOS 通知 (osascript) — 调用系统 AppleScript。"""
# 转义双引号, 避免 AppleScript 注入
safe_title = title.replace('"', '\\"')
safe_msg = message.replace('"', '\\"')
script = (
f'display notification "{safe_msg}" with title "{safe_title}"'
)
result = subprocess.run( # noqa: S603, S607
["osascript", "-e", script],
capture_output=True,
timeout=5,
)
return result.returncode == 0
def _notify_notify_send(title: str, message: str, icon: Path | None) -> bool:
"""Linux 通知 (notify-send) — freedesktop.org 标准。"""
args = ["notify-send", title]
if icon and Path(icon).exists():
args.extend(["--icon", str(icon)])
args.append(message)
result = subprocess.run( # noqa: S603
args,
capture_output=True,
timeout=5,
)
return result.returncode == 0
+250
View File
@@ -0,0 +1,250 @@
"""异步盘后管道任务注册表 — 每个 job 独立 JSON 文件。
设计:
- job_store/ 文件夹,每个 job 一个 {id}.json,最多保留 max_jobs 个文件
- running/pending 状态的 job 仅存内存(高频读写)
- succeeded/failed 后写入独立文件并从内存释放
- 列表查询 = 内存中的活跃 job + 磁盘文件扫描,按时间排序
- 单个查询 = 内存优先,没有则读磁盘
- 创建新 job 前检查文件数量,>= max_jobs 时删除最老的文件
"""
from __future__ import annotations
import json
import logging
import os
import threading
import uuid
from datetime import datetime
from pathlib import Path
from typing import Any, Literal
logger = logging.getLogger(__name__)
JobStatus = Literal["pending", "running", "succeeded", "failed"]
def _default_store_dir() -> Path:
from app.config import settings
return settings.data_dir / "job_store"
_STORE_DIR = _default_store_dir()
class JobStore:
def __init__(self, max_jobs: int = 50, store_dir: Path = _STORE_DIR) -> None:
self._max_jobs = max_jobs
self._store_dir = store_dir
self._active_jobs: dict[str, dict[str, Any]] = {} # running/pending
self._active_id: str | None = None
self._lock = threading.Lock()
self._store_dir.mkdir(parents=True, exist_ok=True)
# ===== persistence =====
def _write_file(self, job: dict[str, Any]) -> None:
"""将终态 job 写入独立 JSON 文件。"""
path = self._store_dir / f"{job['id']}.json"
try:
path.write_text(
json.dumps(job, ensure_ascii=False, indent=None),
encoding="utf-8",
)
except Exception:
logger.warning("failed to write job file %s", path)
def _read_file(self, job_id: str) -> dict[str, Any] | None:
"""从磁盘读取单个 job 文件。"""
path = self._store_dir / f"{job_id}.json"
if not path.exists():
return None
try:
return json.loads(path.read_text("utf-8"))
except Exception:
logger.warning("failed to read job file %s", path)
return None
def _delete_oldest(self) -> None:
"""删除最老的 job 文件,保持文件数量 < max_jobs。"""
try:
files = sorted(self._store_dir.glob("*.json"), key=lambda f: f.stat().st_mtime)
except Exception:
return
while len(files) >= self._max_jobs:
oldest = files.pop(0)
try:
oldest.unlink()
except Exception:
logger.warning("failed to delete old job file %s", oldest)
def _job_files_sorted(self) -> list[dict[str, Any]]:
"""扫描磁盘上所有 job 文件,按 started_at 从新到旧排序。"""
jobs: list[dict[str, Any]] = []
for f in self._store_dir.glob("*.json"):
try:
jobs.append(json.loads(f.read_text("utf-8")))
except Exception:
continue
jobs.sort(key=lambda j: j.get("started_at") or "", reverse=True)
return jobs
# ===== lifecycle =====
def create(self) -> str:
with self._lock:
if self._active_id and self._active_jobs.get(self._active_id, {}).get("status") == "running":
return self._active_id
job_id = uuid.uuid4().hex[:10]
self._active_jobs[job_id] = {
"id": job_id,
"status": "pending",
"stage": "init",
"progress": 0,
"stage_pct": 0,
"log": [],
"started_at": None,
"finished_at": None,
"duration_s": None,
"result": None,
"error": None,
}
self._active_id = job_id
return job_id
def start(self, job_id: str) -> None:
with self._lock:
j = self._active_jobs.get(job_id)
if not j:
return
j["status"] = "running"
j["started_at"] = datetime.utcnow().isoformat(timespec="seconds") + "Z"
def succeed(self, job_id: str, result: Any) -> None:
with self._lock:
j = self._active_jobs.pop(job_id, None)
if not j:
return
j["status"] = "succeeded"
j["finished_at"] = datetime.utcnow().isoformat(timespec="seconds") + "Z"
j["progress"] = 100
j["result"] = result
j["duration_s"] = _duration_s(j)
if self._active_id == job_id:
self._active_id = None
self._delete_oldest()
self._write_file(j)
def fail(self, job_id: str, error: str) -> None:
with self._lock:
j = self._active_jobs.pop(job_id, None)
if not j:
return
j["status"] = "failed"
j["finished_at"] = datetime.utcnow().isoformat(timespec="seconds") + "Z"
j["error"] = error
j["duration_s"] = _duration_s(j)
if self._active_id == job_id:
self._active_id = None
self._delete_oldest()
self._write_file(j)
# ===== progress =====
def progress(self, job_id: str, stage: str, pct: int, msg: str,
stage_pct: int | None = None, skip_log: bool = False) -> None:
with self._lock:
j = self._active_jobs.get(job_id)
if not j:
return
j["stage"] = stage
j["progress"] = max(0, min(100, int(pct)))
if stage_pct is not None:
j["stage_pct"] = max(0, min(100, int(stage_pct)))
elif j["stage"] != stage:
j["stage_pct"] = 0
entry = {
"ts": datetime.utcnow().isoformat(timespec="seconds") + "Z",
"stage": stage,
"msg": msg,
}
if skip_log:
entry["_skip"] = True
if skip_log and j["log"] and j["log"][-1].get("stage") == stage and j["log"][-1].get("_skip"):
j["log"][-1] = entry
else:
j["log"].append(entry)
if len(j["log"]) > 200:
j["log"] = j["log"][-200:]
# ===== query =====
def get(self, job_id: str) -> dict[str, Any] | None:
# 内存中的活跃 job 优先
j = self._active_jobs.get(job_id)
if j:
return j
# 否则从磁盘读
return self._read_file(job_id)
def list_recent(self, limit: int = 20) -> list[dict[str, Any]]:
# 合并: 内存中的活跃 job + 磁盘文件
all_jobs: list[dict[str, Any]] = list(self._active_jobs.values())
all_jobs.extend(self._job_files_sorted())
# 按 started_at 从新到旧排序,去重(理论上不会有重复)
seen: set[str] = set()
result: list[dict[str, Any]] = []
for j in sorted(all_jobs, key=lambda x: x.get("started_at") or "", reverse=True):
jid = j["id"]
if jid in seen:
continue
seen.add(jid)
result.append(_summary(j))
if len(result) >= limit:
break
return result
def active_id(self) -> str | None:
return self._active_id
def clear(self) -> None:
"""清空所有任务(内存 + 磁盘文件)。"""
with self._lock:
self._active_jobs.clear()
self._active_id = None
for f in self._store_dir.glob("*.json"):
try:
f.unlink()
except Exception:
pass
def _summary(j: dict[str, Any]) -> dict[str, Any]:
return {
"id": j["id"],
"status": j["status"],
"stage": j["stage"],
"progress": j["progress"],
"stage_pct": j.get("stage_pct", 0),
"started_at": j["started_at"],
"finished_at": j["finished_at"],
"duration_s": j["duration_s"],
"result": j["result"],
"error": j["error"],
}
def _duration_s(j: dict[str, Any]) -> float | None:
if not j.get("started_at") or not j.get("finished_at"):
return None
try:
s = datetime.fromisoformat(j["started_at"])
e = datetime.fromisoformat(j["finished_at"])
return round((e - s).total_seconds(), 2)
except Exception: # noqa: BLE001
return None
# 进程内单例
job_store = JobStore()
+585
View File
@@ -0,0 +1,585 @@
"""用户偏好设置持久化。
存储位置: data/user_data/preferences.json
沿用 secrets_store 的 merge-write 模式,但不做 chmod 0600 (非敏感数据)。
"""
from __future__ import annotations
import json
import logging
from pathlib import Path
logger = logging.getLogger(__name__)
def _path() -> Path:
from app.config import settings
p = settings.data_dir / "user_data" / "preferences.json"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def load() -> dict:
p = _path()
if p.exists():
try:
return json.loads(p.read_text(encoding="utf-8"))
except Exception as e: # noqa: BLE001
logger.warning("preferences.json malformed: %s", e)
return {}
def save(updates: dict) -> dict:
"""合并写入。返回新内容。"""
current = load()
current.update(updates)
_path().write_text(
json.dumps(current, indent=2, ensure_ascii=False), encoding="utf-8",
)
return current
def get_realtime_quotes_enabled() -> bool:
return load().get("realtime_quotes_enabled", False)
def get_indices_nav_pinned() -> bool:
"""侧栏指数报价卡片是否固定显示。默认 True(常驻)。
关闭后,卡片跟随实时行情开关(仅实时开时显示)。"""
return load().get("indices_nav_pinned", True)
def get_realtime_quote_interval() -> float:
return load().get("realtime_quote_interval", 10.0)
def get_realtime_watchlist_symbols() -> list[str]:
"""Free 档自选实时监控标的:直接取自选页前 5 个。"""
try:
from app.services import watchlist
rows = watchlist.list_symbols()
except Exception as e: # noqa: BLE001
logger.warning("load watchlist for realtime failed: %s", e)
return []
out: list[str] = []
for row in rows:
symbol = str((row or {}).get("symbol") or "").strip().upper()
if symbol and symbol not in out:
out.append(symbol)
if len(out) >= 5:
break
return out
def set_realtime_watchlist_symbols(symbols: list[str]) -> list[str]: # noqa: ARG001
"""兼容旧接口: Free 实时标的现在由自选页前 5 个决定。"""
return get_realtime_watchlist_symbols()
def set_realtime_quote_interval(interval: float) -> float:
"""保存行情轮询间隔(不在此做 min/max 校验,由调用方按档位限制)。"""
current = load()
current["realtime_quote_interval"] = interval
_path().write_text(
json.dumps(current, indent=2, ensure_ascii=False), encoding="utf-8",
)
return interval
def get_minute_sync_enabled() -> bool:
return load().get("minute_sync_enabled", False)
def get_minute_sync_days() -> int:
return max(1, min(30, load().get("minute_sync_days", 5)))
# ===== 数据源选择 (默认 TickFlow;第一阶段仅日K切换入口) =====
_ALLOWED_DATA_PROVIDERS = {"tickflow"}
def get_daily_data_provider() -> str:
provider = str(load().get("daily_data_provider", "tickflow") or "tickflow").lower()
return provider if provider in _ALLOWED_DATA_PROVIDERS else "tickflow"
def get_adj_factor_provider() -> str:
provider = str(load().get("adj_factor_provider", "same_as_daily") or "same_as_daily").lower()
if provider == "same_as_daily":
return provider
return provider if provider in _ALLOWED_DATA_PROVIDERS else "same_as_daily"
def get_minute_data_provider() -> str:
provider = str(load().get("minute_data_provider", "tickflow") or "tickflow").lower()
return provider if provider in _ALLOWED_DATA_PROVIDERS else "tickflow"
def get_realtime_data_provider() -> str:
# 盘中实时现阶段仅支持 TickFlow。
return "tickflow"
# ===== 盘后管道拉取内容开关 (A股 / ETF / 指数 独立控制) =====
def get_pipeline_pull_a_share() -> bool:
"""A 股日K固定拉取。"""
return True
def get_pipeline_pull_etf() -> bool:
"""是否拉取 ETF 日K。默认 False(标的多,首次较慢)。"""
return load().get("pipeline_pull_etf", False)
def get_pipeline_pull_index() -> bool:
"""是否拉取指数日K。默认 True。"""
return load().get("pipeline_pull_index", True)
_PIPELINE_PULL_KEYS = ("pipeline_pull_etf", "pipeline_pull_index")
def get_pipeline_pull_types() -> dict:
"""返回三个拉取开关的当前值。"""
return {
"pipeline_pull_a_share": get_pipeline_pull_a_share(),
"pipeline_pull_etf": get_pipeline_pull_etf(),
"pipeline_pull_index": get_pipeline_pull_index(),
}
def set_pipeline_pull_types(cfg: dict) -> dict:
"""批量保存拉取开关。只接受白名单内的布尔字段。"""
updates = {
k: bool(v) for k, v in cfg.items()
if k in _PIPELINE_PULL_KEYS and v is not None
}
save(updates)
return get_pipeline_pull_types()
def get_pipeline_index_symbols() -> str:
"""指数自定义拉取代码(逗号/换行/空格分隔)。空串表示全量。"""
return str(load().get("pipeline_index_symbols", "") or "").strip()
def set_pipeline_index_symbols(symbols: str) -> str:
"""保存指数自定义代码,返回规范化后的字符串。"""
save({"pipeline_index_symbols": symbols})
return get_pipeline_index_symbols()
def get_pipeline_schedule() -> dict:
"""返回盘后管道调度时间 {"hour": 15, "minute": 30}。"""
d = load().get("pipeline_schedule", {"hour": 15, "minute": 30})
return {"hour": d.get("hour", 15), "minute": d.get("minute", 30)}
def set_pipeline_schedule(hour: int, minute: int) -> dict:
h = max(0, min(23, hour))
m = max(0, min(59, minute))
# 盘后不早于 15:00
if h * 60 + m < 15 * 60:
h, m = 15, 0
save({"pipeline_schedule": {"hour": h, "minute": m}})
return {"hour": h, "minute": m}
def get_instruments_schedule() -> dict:
"""返回盘前标的维表调度时间 {"hour": 9, "minute": 10}。"""
d = load().get("instruments_schedule", {"hour": 9, "minute": 10})
return {"hour": d.get("hour", 9), "minute": d.get("minute", 10)}
def set_instruments_schedule(hour: int, minute: int) -> dict:
h = max(0, min(23, hour))
m = max(0, min(59, minute))
# 盘前不晚于 09:15
if h * 60 + m > 9 * 60 + 15:
h, m = 9, 15
save({"instruments_schedule": {"hour": h, "minute": m}})
return {"hour": h, "minute": m}
def get_enriched_batch_size() -> int:
"""返回 enriched 全量计算每批 symbol 数量。"""
return max(1, min(10000, load().get("enriched_batch_size", 1000)))
def set_enriched_batch_size(size: int) -> int:
"""保存 enriched 全量计算批次大小。"""
size = max(10, min(6000, size))
save({"enriched_batch_size": size})
return size
def get_index_daily_batch_size() -> int:
"""返回指数日 K 同步每批 symbol 数量。"""
return max(1, min(10000, load().get("index_daily_batch_size", 100)))
def set_index_daily_batch_size(size: int) -> int:
"""保存指数日 K 同步批次大小。"""
size = max(1, min(10000, size))
save({"index_daily_batch_size": size})
return size
# ── 五档盘口 sealed(真假涨停) 配置 ──────────────────────
def get_limit_ladder_monitor_enabled() -> bool:
"""连板梯队 5 档监控开关。关闭时 depth 不轮询(连板梯队降级显示)。"""
return load().get("limit_ladder_monitor_enabled", False)
def get_depth_polling_interval() -> float:
"""depth 盘中轮询间隔(秒)。默认 20(Pro/Expert 都适用)。"""
return float(load().get("depth_polling_interval", 20.0))
def set_depth_polling_interval(interval: float) -> float:
"""保存 depth 轮询间隔。套餐范围 clamp 由 depth_service 按档位做。"""
interval = max(1.0, min(600.0, float(interval)))
save({"depth_polling_interval": interval})
return interval
def get_depth_finalize_time() -> dict:
"""盘后 sealed 定版时间 {"hour": 15, "minute": 2}。范围 15:01~18:00。"""
d = load().get("depth_finalize_time", {"hour": 15, "minute": 2})
return {"hour": d.get("hour", 15), "minute": d.get("minute", 2)}
def set_depth_finalize_time(hour: int, minute: int) -> dict:
"""保存盘后 sealed 定版时间,强制范围 15:01~18:00。"""
h = max(0, min(23, hour))
m = max(0, min(59, minute))
# 下限 15:01, 上限 18:00
if h * 60 + m < 15 * 60 + 1:
h, m = 15, 1
if h * 60 + m > 18 * 60:
h, m = 18, 0
save({"depth_finalize_time": {"hour": h, "minute": m}})
return {"hour": h, "minute": m}
# 复盘推送可选渠道白名单 (微信等暂未实现, 不在白名单内, 前端仅作占位)
# 多选: 不推送 = 空数组, 而非 'none'
REVIEW_PUSH_CHANNELS = {"feishu"}
def get_review_schedule() -> dict:
"""定时复盘调度 {"enabled": False, "hour": 15, "minute": 10}。默认关闭。
A股 15:00 收盘, 默认时间设为 15:10(收盘后即时复盘), 强制下限 15:00。
"""
d = load().get("review_schedule", {"enabled": False, "hour": 15, "minute": 10})
return {
"enabled": bool(d.get("enabled", False)),
"hour": d.get("hour", 15),
"minute": d.get("minute", 10),
}
def set_review_schedule(enabled: bool, hour: int, minute: int) -> dict:
"""保存定时复盘调度。强制时间下限 15:00(A股收盘)。
enabled=False 时时间仍保存(下次开启可沿用), 但调度器不会注册 job。
"""
h = max(0, min(23, hour))
m = max(0, min(59, minute))
# 下限 15:00: A股 15:00 收盘, 收盘后才有当日完整数据复盘
if h * 60 + m < 15 * 60:
h, m = 15, 0
save({"review_schedule": {"enabled": bool(enabled), "hour": h, "minute": m}})
return {"enabled": bool(enabled), "hour": h, "minute": m}
def get_review_push_channels() -> list[str]:
"""复盘推送渠道(多选) — 选定的外部工具列表, 复盘归档后逐个推送。
与 review_schedule / 实时行情完全独立, 常驻可单独设置。
空列表 = 不推送; ['feishu'] = 推送到飞书(复用监控中心全局 feishu_webhook_url/secret)。
向后兼容:
- 老多版本单选 review_push_channel=='feishu' → ['feishu']
- 更老布尔 review_push_enabled==True → ['feishu']
"""
d = load()
raw = d.get("review_push_channels")
if isinstance(raw, list):
return [c for c in raw if c in REVIEW_PUSH_CHANNELS]
# 兼容老单选字符串
if d.get("review_push_channel") == "feishu":
return ["feishu"]
# 兼容更老布尔开关
if d.get("review_push_enabled") is True:
return ["feishu"]
return []
def set_review_push_channels(channels: list[str]) -> list[str]:
"""保存复盘推送渠道(多选)。过滤白名单外的值、去重、保序。空列表 = 不推送。"""
seen: set[str] = set()
cleaned: list[str] = []
for c in channels or []:
if c in REVIEW_PUSH_CHANNELS and c not in seen:
seen.add(c)
cleaned.append(c)
save({"review_push_channels": cleaned})
return cleaned
# ===== 实时监控 =====
# 页面 SSE 刷新配置: { "watchlist": true, "monitor": true, ... }
# 可刷新的页面列表及其默认值
SSE_REFRESH_PAGES_DEFAULT = {
"watchlist": True,
"limit-ladder": False,
}
SIDEBAR_INDEX_SYMBOLS_DEFAULT = ["000001.SH", "399001.SZ", "399006.SZ", "000680.SH"]
# ===== 盘中实时行情范围 (独立于盘后管道范围) =====
def get_realtime_pull_stock() -> bool:
return load().get("realtime_pull_stock", True)
def get_realtime_pull_etf() -> bool:
# 老用户兼容: ETF 实时默认关闭,避免升级后请求量/写盘量突然增加。
return load().get("realtime_pull_etf", False)
def get_realtime_pull_index() -> bool:
return load().get("realtime_pull_index", True)
def get_realtime_index_mode() -> str:
mode = str(load().get("realtime_index_mode", "core") or "core").lower()
return mode if mode in {"core", "all"} else "core"
def get_realtime_index_symbols() -> list[str]:
stored = load().get("realtime_index_symbols", SIDEBAR_INDEX_SYMBOLS_DEFAULT)
if isinstance(stored, str):
import re
stored = [s.strip() for s in re.split(r"[,\s]+", stored) if s.strip()]
return [str(s) for s in stored if str(s).strip()]
def set_realtime_quote_scope(cfg: dict) -> dict:
updates = {}
for key in ("realtime_pull_stock", "realtime_pull_etf", "realtime_pull_index"):
if key in cfg and cfg[key] is not None:
updates[key] = bool(cfg[key])
if "realtime_index_mode" in cfg and cfg["realtime_index_mode"] in {"core", "all"}:
updates["realtime_index_mode"] = cfg["realtime_index_mode"]
if "realtime_index_symbols" in cfg and cfg["realtime_index_symbols"] is not None:
updates["realtime_index_symbols"] = cfg["realtime_index_symbols"]
if updates:
save(updates)
return get_realtime_quote_scope()
def get_realtime_quote_scope() -> dict:
return {
"realtime_pull_stock": get_realtime_pull_stock(),
"realtime_pull_etf": get_realtime_pull_etf(),
"realtime_pull_index": get_realtime_pull_index(),
"realtime_index_mode": get_realtime_index_mode(),
"realtime_index_symbols": get_realtime_index_symbols(),
}
def get_sse_refresh_pages() -> dict[str, bool]:
"""返回每个页面的 SSE 刷新开关。"""
stored = load().get("sse_refresh_pages", {})
# 合并默认值 (新增页面自动出现)
result = dict(SSE_REFRESH_PAGES_DEFAULT)
result.update(stored)
return result
def set_sse_refresh_pages(pages: dict[str, bool]) -> dict[str, bool]:
"""保存页面 SSE 刷新配置。"""
save({"sse_refresh_pages": pages})
return get_sse_refresh_pages()
def get_sidebar_index_symbols() -> list[str]:
"""返回左侧菜单显示的指数代码。"""
stored = load().get("sidebar_index_symbols", SIDEBAR_INDEX_SYMBOLS_DEFAULT)
allowed = set(SIDEBAR_INDEX_SYMBOLS_DEFAULT)
return [s for s in stored if s in allowed]
def get_strategy_monitor_enabled() -> bool:
"""策略告警评估总开关。"""
return load().get("strategy_monitor_enabled", False)
def get_system_notify_enabled() -> bool:
"""系统通知开关 — 开启后监控告警同时推送到操作系统通知中心。"""
return load().get("system_notify_enabled", False)
def set_system_notify_enabled(enabled: bool) -> bool:
"""保存系统通知开关。"""
save({"system_notify_enabled": bool(enabled)})
return bool(enabled)
def get_feishu_webhook_url() -> str:
"""飞书自定义机器人 Webhook 地址 — 全局共用一处, 所有启用推送的规则都推到这一个群。"""
return load().get("feishu_webhook_url", "")
def get_feishu_webhook_secret() -> str:
"""飞书自定义机器人签名密钥 — 机器人启用「签名校验」时必填, 留空表示不验签。"""
return load().get("feishu_webhook_secret", "")
def set_feishu_webhook_url(url: str) -> str:
"""保存飞书 Webhook 地址。传入空串表示清空配置。"""
save({"feishu_webhook_url": str(url or "").strip()})
return get_feishu_webhook_url()
def set_feishu_webhook_secret(secret: str) -> str:
"""保存飞书签名密钥。传入空串表示不验签。"""
save({"feishu_webhook_secret": str(secret or "").strip()})
return get_feishu_webhook_secret()
def get_webhook_enabled_default() -> bool:
"""新建监控规则时是否默认勾选「飞书推送」。
数据模型当前只有一个 webhook_enabled 布尔 (即飞书), QMT/ptrade 待定。
此默认值供规则编辑器新建规则时预填, 单条规则仍可独立修改。
"""
return load().get("webhook_enabled_default", False)
def set_webhook_enabled_default(enabled: bool) -> bool:
"""保存飞书推送默认勾选态。"""
save({"webhook_enabled_default": bool(enabled)})
return get_webhook_enabled_default()
def get_screener_auto_run() -> bool:
"""选股页进入时是否自动运行所有策略 (获取命中数)。默认开。"""
return load().get("screener_auto_run", True)
def get_strategy_monitor_ids() -> list[str]:
"""返回监控池中的策略 ID。"""
return load().get("strategy_monitor_ids", [])
def set_realtime_monitor_config(cfg: dict) -> dict:
"""批量更新实时监控配置。"""
updates = {}
if "sse_refresh_pages" in cfg:
updates["sse_refresh_pages"] = cfg["sse_refresh_pages"]
if "strategy_monitor_enabled" in cfg:
updates["strategy_monitor_enabled"] = cfg["strategy_monitor_enabled"]
if "strategy_monitor_ids" in cfg:
updates["strategy_monitor_ids"] = cfg["strategy_monitor_ids"]
if "sidebar_index_symbols" in cfg:
allowed = set(SIDEBAR_INDEX_SYMBOLS_DEFAULT)
updates["sidebar_index_symbols"] = [s for s in cfg["sidebar_index_symbols"] if s in allowed]
if "screener_auto_run" in cfg:
updates["screener_auto_run"] = bool(cfg["screener_auto_run"])
if updates:
save(updates)
return get_realtime_monitor_config()
def get_realtime_monitor_config() -> dict:
"""返回完整的实时监控配置。"""
return {
"sse_refresh_pages": get_sse_refresh_pages(),
"strategy_monitor_enabled": get_strategy_monitor_enabled(),
"strategy_monitor_ids": get_strategy_monitor_ids(),
"sidebar_index_symbols": get_sidebar_index_symbols(),
"screener_auto_run": get_screener_auto_run(),
}
def get_nav_order() -> list[str]:
"""返回左侧菜单的自定义排序(内置页面 path + 扩展分析菜单 id)。"""
return load().get("nav_order", [])
def set_nav_order(order: list[str]) -> list[str]:
"""保存左侧菜单排序。"""
save({"nav_order": order})
return get_nav_order()
def get_nav_hidden() -> list[str]:
"""返回左侧菜单中隐藏的项 id 列表。"""
return load().get("nav_hidden", [])
def set_nav_hidden(hidden: list[str]) -> list[str]:
"""保存左侧菜单隐藏项。"""
save({"nav_hidden": hidden})
return get_nav_hidden()
def get_watchlist_columns() -> list[dict] | None:
"""返回自选列表列配置。"""
return load().get("watchlist_columns")
def set_watchlist_columns(columns: list[dict]) -> list[dict]:
"""保存自选列表列配置。"""
save({"watchlist_columns": columns})
return columns
def get_screener_result_columns() -> list[dict] | None:
"""返回策略结果列表列配置。"""
return load().get("screener_result_columns")
def set_screener_result_columns(columns: list[dict]) -> list[dict]:
"""保存策略结果列表列配置。"""
save({"screener_result_columns": columns})
return columns
# ===== 首次使用引导 =====
def get_onboarding_completed() -> bool:
"""是否已完成首次使用向导。默认 False(新用户)。"""
return bool(load().get("onboarding_completed", False))
def set_onboarding_completed(done: bool = True) -> bool:
"""标记首次使用向导完成状态。"""
save({"onboarding_completed": bool(done)})
return bool(done)
# ===== 财务数据同步时间(持久化,重启不丢失) =====
# 结构: { "metrics": "2026-06-25T10:00:00+08:00", "income": ..., ... }
def get_financial_sync_times() -> dict[str, str]:
"""返回各财务表的最后同步时间(ISO 字符串)。未同步过的表不在返回值中。"""
return load().get("financial_sync_times", {}) or {}
def set_financial_sync_time(table: str, iso_ts: str) -> None:
"""更新单张财务表的最后同步时间(合并写入,不清除其他表)。"""
times = get_financial_sync_times()
times[table] = iso_ts
save({"financial_sync_times": times})
+971
View File
@@ -0,0 +1,971 @@
"""全局实时行情服务。
集中管理全市场行情拉取 + enriched 缓存,供盘中选股、自选股等所有模块复用。
架构:
- 后台线程轮询 TickFlow get_by_universes(["CN_Equity_A", "CN_Index"])
- 拉取行情 → 写 kline_daily (不复权) + 增量计算 enriched → 写盘 + 更新缓存
- _enriched_cache 是唯一的盘中数据源 (OHLCV + 全套技术指标)
- _live_agg_cache 是递推状态 (只加载一次, 盘中不变)
数据流 (每轮 ~15s):
1. API 拉取 → raw_records (临时变量)
2. raw_records → 写 kline_daily (不复权原始价格)
3. raw_records → 更新 _enriched_cache 的 OHLCV
4. 增量计算 enriched 指标 (~50ms)
5. 写 kline_daily_enriched + 替换 _enriched_cache
6. 通知 SSE
生命周期:
- 服务启动时读取 preferences,若 enabled 则自动启动线程
- 运行中可通过 API 切换开关
- 关闭时停止线程
"""
from __future__ import annotations
import logging
import threading
import time
from datetime import date, datetime, time as dt_time
import polars as pl
logger = logging.getLogger(__name__)
class QuoteService:
"""全局实时行情服务 — 单例。"""
CORE_INDEX_SYMBOLS = ("000001.SH", "399001.SZ", "399006.SZ", "000680.SH")
# 档位 → 最小轮询间隔 (秒)
TIER_MIN_INTERVAL = {
"expert": 1.0,
"pro": 2.0,
"starter": 3.0,
"free": 6.0,
}
DEFAULT_INTERVAL = 10.0
MAX_INTERVAL = 60.0
def __init__(self) -> None:
self._lock = threading.Lock()
self._running = False
self._enabled = False # 全局开关 (持久化到 preferences)
self._interval = self.DEFAULT_INTERVAL
self._thread: threading.Thread | None = None
self._repo = None # 延迟注入, 避免循环导入
self._update_event = threading.Event() # SSE 通知: 行情更新后 set
self._alert_event = threading.Event() # SSE 通知: 有告警时 set
self._depth_update_event = threading.Event() # SSE 通知: depth 五档修正后 set (刷新连板梯队)
self._pending_alerts: list[dict] = [] # 待推送的告警
self._max_pending_alerts: int = 1000 # 背压上限: 超出丢弃最旧
# 复盘进度 SSE 通道: 定时复盘流式生成时, 把 meta/delta/done 事件推给开着页面的前端
self._review_event = threading.Event() # SSE 通知: 有复盘进度事件时 set
self._pending_review: list[str] = [] # 待推送的复盘事件(JSON 字符串)
self._max_pending_review: int = 200 # 背压上限: 超出丢弃最旧
self._strategy_monitor = None # 延迟注入
self._app_state = None # 延迟注入 (FastAPI app.state)
# 拉取元信息 (给 SSE / status 用)
self._fetch_time: float = 0.0 # perf_counter (用于计算 quote_age_ms)
self._fetch_ms: float = 0.0 # 拉取耗时 (毫秒)
self._fetched_at: float = 0.0 # 拉取完成的 Unix 时间戳 (毫秒)
self._symbol_count: int = 0
self._index_symbol_count: int = 0
self._etf_symbol_count: int = 0
self._index_quotes_cache: pl.DataFrame | None = None
# ================================================================
# 生命周期
# ================================================================
def start(self, interval: float = 0.0) -> None:
"""启动后台行情轮询线程。"""
if self._running:
return
if interval <= 0:
from app.services import preferences
interval = preferences.get_realtime_quote_interval()
self._interval = self._clamp_interval(interval)
self._running = True
self._enabled = True
self._thread = threading.Thread(target=self._poll_loop, daemon=True)
self._thread.start()
self._save_enabled(True)
logger.info("行情服务已启动, 轮询间隔 %.1fs", self._interval)
def stop(self) -> None:
"""停止后台行情轮询线程。"""
self._running = False
self._enabled = False
if self._thread:
self._thread.join(timeout=10)
self._thread = None
self._save_enabled(False)
logger.info("行情服务已停止")
def enable(self) -> bool:
"""开启自动行情 (不立即启动线程,等下一个交易时段)。
none 档无实时行情权限,拒绝开启并返回 False;
free 档开启自选股实时,starter+ 开启全市场实时。返回值表示是否真正开启。
"""
if not self.is_realtime_allowed():
logger.warning("实时行情开启被拒:当前档位(none)无实时行情权限")
return False
self._enabled = True
self._save_enabled(True)
if not self._running:
from app.services import preferences
self._interval = self._clamp_interval(preferences.get_realtime_quote_interval())
self._running = True
self._thread = threading.Thread(target=self._poll_loop, daemon=True)
self._thread.start()
logger.info("行情服务已启用, 轮询间隔 %.1fs", self._interval)
def disable(self) -> None:
"""关闭自动行情。"""
self.stop()
logger.info("行情服务已关闭")
def boot_check(self) -> None:
"""启动时检查 preferences,若 enabled 则自动启动。
none 档无实时行情权限:即使 preferences 标记为 enabled,
也不启动,并同步 preferences 为关闭(避免 UI 误显示已开启)。
"""
from app.services import preferences
if not self.is_realtime_allowed():
if preferences.get_realtime_quotes_enabled():
self._save_enabled(False)
logger.info("实时行情未启动:当前档位(none)无实时行情权限")
return
if preferences.get_realtime_quotes_enabled():
self.start()
def set_repo(self, repo) -> None:
"""注入 KlineRepository, 用于实时落盘。"""
self._repo = repo
def set_app_state(self, app_state) -> None:
"""注入 FastAPI app.state, 用于获取 strategy_monitor 等单例。"""
self._app_state = app_state
def set_interval(self, interval: float) -> float:
"""运行时更新轮询间隔(立即生效)。"""
clamped = self._clamp_interval(interval)
self._interval = clamped
from app.services import preferences
preferences.set_realtime_quote_interval(clamped)
logger.info("轮询间隔已更新为 %.1fs", clamped)
return clamped
def get_min_interval(self) -> float:
"""返回当前档位允许的最小间隔。"""
return self._tier_min_interval()
def wait_for_update(self, timeout: float = 30.0) -> bool:
"""阻塞等待下一次行情更新 (供 SSE 线程使用)。"""
self._update_event.clear()
return self._update_event.wait(timeout=timeout)
def wait_for_alert(self, timeout: float = 30.0) -> bool:
"""阻塞等待告警 (供 SSE 线程使用)。"""
self._alert_event.clear()
return self._alert_event.wait(timeout=timeout)
def notify_depth_updated(self) -> None:
"""五档盘口修正完成后调用: 通知 SSE 推送 depth_updated, 触发连板梯队刷新。
与行情/告警通道独立 — 只刷新连板梯队, 不连带刷新 watchlist 等。
"""
self._depth_update_event.set()
def wait_for_depth_update(self, timeout: float = 30.0) -> bool:
"""阻塞等待 depth 修正 (供 SSE 线程使用)。"""
self._depth_update_event.clear()
return self._depth_update_event.wait(timeout=timeout)
def pop_alerts(self) -> list[dict]:
"""取走所有待推送的告警 (线程安全)。"""
with self._lock:
alerts = self._pending_alerts
self._pending_alerts = []
return alerts
# ================================================================
# 复盘进度 SSE 通道 — 定时复盘流式生成时, 把事件实时推给前端
# ================================================================
def push_review_event(self, event_json: str) -> None:
"""追加一条复盘进度事件(JSON 字符串), 并唤醒 SSE generator。
事件格式与 recap_market_stream 的产出一致(meta/delta/error/done),
前端 reviewStore 直接消费。背压: 超过上限丢弃最旧(复盘流几百条 delta, 200 够用)。
"""
with self._lock:
self._pending_review.append(event_json)
if len(self._pending_review) > self._max_pending_review:
overflow = len(self._pending_review) - self._max_pending_review
self._pending_review = self._pending_review[overflow:]
self._review_event.set()
def wait_for_review(self, timeout: float = 30.0) -> bool:
"""阻塞等待复盘进度事件 (供 SSE 线程使用)。"""
self._review_event.clear()
return self._review_event.wait(timeout=timeout)
def pop_review_events(self) -> list[str]:
"""取走所有待推送的复盘事件 (线程安全)。"""
with self._lock:
events = self._pending_review
self._pending_review = []
return events
# ================================================================
# 档位感知间隔限制
# ================================================================
@staticmethod
def _current_tier() -> str:
"""获取当前档位名(小写)。"""
from app.tickflow.policy import tier_label
return tier_label().split()[0].split("+")[0].strip().lower()
@classmethod
def realtime_mode(cls) -> str:
"""当前实时行情模式: none / watchlist / full_market。"""
tier = cls._current_tier()
if tier == "none":
return "none"
if tier == "free":
return "watchlist"
return "full_market"
@classmethod
def is_realtime_allowed(cls) -> bool:
"""当前档位是否允许使用实时行情。"""
return cls.realtime_mode() != "none"
@classmethod
def _tier_min_interval(cls) -> float:
tier = cls._current_tier()
return cls.TIER_MIN_INTERVAL.get(tier, cls.DEFAULT_INTERVAL)
def _clamp_interval(self, interval: float) -> float:
return max(self._tier_min_interval(), min(self.MAX_INTERVAL, interval))
# ================================================================
# 行情数据访问
# ================================================================
def get_enriched_today(self) -> tuple[pl.DataFrame, date | None]:
"""返回今天 enriched 数据 + 日期 (线程安全)。
所有页面统一通过此方法获取实时行情 + 技术指标。
"""
if not self._repo:
return pl.DataFrame(), None
return self._repo.get_enriched_latest()
def get_quotes_compat(self) -> pl.DataFrame:
"""兼容接口: 返回行情 DataFrame (用于盘中选股等需要 last_price/prev_close 的场景)。
从 _enriched_cache 取 today 的数据, 只选行情基础列, 补上 last_price 别名。
不返回指标列, 避免 JOIN live_agg 时列名冲突。
"""
df, _ = self.get_enriched_today()
if df.is_empty():
return df
# 只取盘中选股需要的行情基础列
keep = [c for c in [
"symbol", "close", "open", "high", "low", "volume", "amount",
"prev_close", "change_pct", "change_amount", "amplitude", "turnover_rate",
] if c in df.columns]
df = df.select(keep)
# enriched 的 close 等价于 last_price
if "close" in df.columns and "last_price" not in df.columns:
df = df.with_columns(pl.col("close").alias("last_price"))
return df
def get_index_quotes(self, symbols: list[str] | None = None) -> pl.DataFrame:
"""返回实时指数行情缓存。不会触发 TickFlow 请求。"""
with self._lock:
df = self._index_quotes_cache.clone() if self._index_quotes_cache is not None else pl.DataFrame()
if df.is_empty():
return df
if symbols:
return df.filter(pl.col("symbol").is_in(symbols))
return df
def status(self) -> dict:
"""返回行情服务状态。"""
from app.services import preferences
age = (time.perf_counter() - self._fetch_time) * 1000 if self._fetch_time else -1
mode = self.realtime_mode()
return {
"enabled": self._enabled,
"running": self._running,
"mode": mode,
"realtime_allowed": mode != "none",
"watchlist_symbol_count": len(preferences.get_realtime_watchlist_symbols()),
"interval_s": self._interval,
"symbol_count": self._symbol_count,
"index_symbol_count": self._index_symbol_count,
"etf_symbol_count": self._etf_symbol_count,
"quote_age_ms": round(age, 0) if age >= 0 else None,
"is_trading_hours": self._is_trading_hours(),
"last_fetch_ms": round(self._fetched_at, 0) if self._fetched_at else None,
}
def refresh(self) -> dict:
"""手动触发一次行情拉取。"""
self._fetch_quotes()
return self.status()
# ================================================================
# 后台轮询
# ================================================================
def _poll_loop(self) -> None:
while self._running and self._enabled:
try:
if self._is_trading_hours():
self._fetch_quotes()
else:
logger.debug("非交易时段, 跳过行情轮询")
except Exception as e: # noqa: BLE001
logger.warning("行情轮询异常: %s", e)
waited = 0.0
while self._running and self._enabled and waited < self._interval:
time.sleep(0.5)
waited += 0.5
def _fetch_quotes(self) -> None:
"""按当前档位拉取行情。"""
if self.realtime_mode() == "watchlist":
self._fetch_watchlist_quotes()
return
self._fetch_full_market_quotes()
def _fetch_full_market_quotes(self) -> None:
"""拉取全市场行情 → 写 daily + 计算 enriched + 更新缓存。"""
from app.tickflow.client import get_paid_realtime_client
tf = get_paid_realtime_client()
if tf is None:
logger.warning("实时行情拉取失败:未配置付费服务器 API Key")
return
t0 = time.perf_counter()
now_ts = time.perf_counter()
try:
from app.services import preferences
all_index_symbols = set(self._repo.get_index_symbol_set()) if self._repo else set()
core_index_symbols = set(preferences.get_realtime_index_symbols() or self.CORE_INDEX_SYMBOLS)
all_index_symbols.update(core_index_symbols)
all_etf_symbols = set()
if self._repo:
etf_inst = self._repo.get_etf_instruments()
if not etf_inst.is_empty() and "symbol" in etf_inst.columns:
all_etf_symbols = set(etf_inst["symbol"].cast(pl.Utf8).to_list())
universes: list[str] = []
if preferences.get_realtime_pull_stock():
universes.append("CN_Equity_A")
if preferences.get_realtime_pull_etf() and all_etf_symbols:
universes.append("CN_ETF")
if preferences.get_realtime_pull_index() and preferences.get_realtime_index_mode() == "all":
universes.append("CN_Index")
resp = []
if universes:
resp.extend(tf.quotes.get_by_universes(universes=universes) or [])
if preferences.get_realtime_pull_index() and preferences.get_realtime_index_mode() == "core":
resp.extend(tf.quotes.get(symbols=sorted(core_index_symbols)) or [])
except Exception as e: # noqa: BLE001
logger.warning("行情拉取失败: %s", e)
return
if not resp:
logger.warning("行情数据为空")
return
# ---- 解析 API 响应 (临时变量, 用完丢弃) ----
records = []
for q in resp:
ext = q.get("ext") or {}
last_price = q.get("last_price")
prev_close = q.get("prev_close")
change_amount = ext.get("change_amount")
change_pct = ext.get("change_pct")
if change_amount is None and last_price is not None and prev_close is not None:
change_amount = float(last_price) - float(prev_close)
if change_pct is None and change_amount is not None and prev_close not in (None, 0):
change_pct = float(change_amount) / float(prev_close) * 100
records.append({
"symbol": q.get("symbol"),
"name": q.get("name") or ext.get("name"),
"last_price": last_price,
"prev_close": prev_close,
"open": q.get("open"),
"high": q.get("high"),
"low": q.get("low"),
"volume": q.get("volume"),
"amount": q.get("amount"),
"change_pct": change_pct,
"change_amount": change_amount,
"amplitude": ext.get("amplitude"),
"turnover_rate": ext.get("turnover_rate"),
"timestamp": q.get("timestamp"),
"session": q.get("session"),
})
index_records = [r for r in records if r.get("symbol") in all_index_symbols]
etf_records = [r for r in records if r.get("symbol") in all_etf_symbols]
stock_records = [
r for r in records
if r.get("symbol") not in all_index_symbols and r.get("symbol") not in all_etf_symbols
]
fetch_ms = (time.perf_counter() - t0) * 1000
fetched_at = time.time() * 1000
# ---- 更新元信息 ----
with self._lock:
self._fetch_time = now_ts
self._fetch_ms = fetch_ms
self._fetched_at = fetched_at
self._symbol_count = len(stock_records)
self._index_symbol_count = len(index_records)
self._etf_symbol_count = len(etf_records)
self._index_quotes_cache = self._build_index_quotes(index_records)
logger.info("行情刷新: %d 只股票, %d 只ETF, %d 只指数, 耗时 %.0fms", len(stock_records), len(etf_records), len(index_records), fetch_ms)
# ---- 写 kline_daily (不复权原始价格, 只有 OHLCV) ----
daily_df = self._build_daily(stock_records)
if not daily_df.is_empty() and self._repo:
try:
self._repo.flush_live_daily(daily_df)
except Exception as e: # noqa: BLE001
logger.warning("日K写盘失败: %s", e)
etf_daily_df = self._build_daily(etf_records)
if not etf_daily_df.is_empty() and self._repo:
try:
self._repo.flush_live_daily_asset("etf", etf_daily_df)
except Exception as e: # noqa: BLE001
logger.warning("ETF 日K写盘失败: %s", e)
# ---- 构建 API 直接值的补充表 (不写 daily, 只用于 enriched 计算) ----
quote_extra = self._build_quote_extra(stock_records)
etf_quote_extra = self._build_quote_extra(etf_records)
# ---- 增量计算 enriched + 写盘 + 更新缓存 ----
if not daily_df.is_empty() and self._repo:
self._flush_live_enriched(daily_df, quote_extra, asset_type="stock")
if not etf_daily_df.is_empty() and self._repo:
self._flush_live_enriched(etf_daily_df, etf_quote_extra, asset_type="etf")
# ---- 通知 SSE ----
self._update_event.set()
# ---- 策略监控 + 告警评估 ----
self._evaluate_monitors(daily_df, quote_extra)
def _fetch_watchlist_quotes(self) -> None:
"""Free 档自选股实时: 只拉取最多 5 个 symbols。"""
from app.services import preferences
from app.tickflow.client import get_paid_realtime_client
symbols = preferences.get_realtime_watchlist_symbols()
if not symbols:
logger.info("自选实时未配置标的, 跳过行情拉取")
return
tf = get_paid_realtime_client()
if tf is None:
logger.warning("自选实时拉取失败:未配置付费服务器 API Key")
return
t0 = time.perf_counter()
now_ts = time.perf_counter()
try:
resp = tf.quotes.get(symbols=symbols) or []
except Exception as e: # noqa: BLE001
logger.warning("自选实时拉取失败: %s", e)
return
if not resp:
logger.warning("自选实时行情数据为空")
return
records = []
for q in resp:
ext = q.get("ext") or {}
last_price = q.get("last_price")
prev_close = q.get("prev_close")
change_amount = ext.get("change_amount")
change_pct = ext.get("change_pct")
if change_amount is None and last_price is not None and prev_close is not None:
change_amount = float(last_price) - float(prev_close)
if change_pct is None and change_amount is not None and prev_close not in (None, 0):
change_pct = float(change_amount) / float(prev_close) * 100
records.append({
"symbol": q.get("symbol"),
"name": q.get("name") or ext.get("name"),
"last_price": last_price,
"prev_close": prev_close,
"open": q.get("open"),
"high": q.get("high"),
"low": q.get("low"),
"volume": q.get("volume"),
"amount": q.get("amount"),
"change_pct": change_pct,
"change_amount": change_amount,
"amplitude": ext.get("amplitude"),
"turnover_rate": ext.get("turnover_rate"),
"timestamp": q.get("timestamp"),
"session": q.get("session"),
})
fetch_ms = (time.perf_counter() - t0) * 1000
fetched_at = time.time() * 1000
with self._lock:
self._fetch_time = now_ts
self._fetch_ms = fetch_ms
self._fetched_at = fetched_at
self._symbol_count = len(records)
self._index_symbol_count = 0
self._etf_symbol_count = 0
self._index_quotes_cache = None
logger.info("自选实时刷新: %d 只股票, 耗时 %.0fms", len(records), fetch_ms)
daily_df = self._build_daily(records)
quote_extra = self._build_quote_extra(records)
if not daily_df.is_empty() and self._repo:
try:
self._repo.merge_live_daily_asset("stock", daily_df)
except Exception as e: # noqa: BLE001
logger.warning("自选实时日K写盘失败: %s", e)
self._flush_live_enriched(daily_df, quote_extra, asset_type="stock", merge=True)
self._update_event.set()
self._evaluate_monitors(daily_df, quote_extra)
# ================================================================
# 工具
# ================================================================
@staticmethod
def _build_daily(records: list[dict]) -> pl.DataFrame:
"""将 API records 转为日K格式 DataFrame (只有 OHLCV, 写 kline_daily 用)。"""
if not records:
return pl.DataFrame()
df = pl.DataFrame(records)
cols_map = {
"symbol": "symbol",
"last_price": "close",
"open": "open",
"high": "high",
"low": "low",
"volume": "volume",
"amount": "amount",
}
select_exprs = []
for src, dst in cols_map.items():
if src in df.columns:
select_exprs.append(pl.col(src).alias(dst))
if not select_exprs:
return pl.DataFrame()
result = df.select(select_exprs).with_columns(
pl.lit(date.today()).cast(pl.Date).alias("date"),
)
# 修复: API 在非交易时段可能返回 open/high/low=0 或 null,
# 导致蜡烛从 0 开始。用 close 填充这些异常值。
for col in ("open", "high", "low"):
if col in result.columns:
result = result.with_columns(
pl.when((pl.col(col) == 0) | pl.col(col).is_null())
.then(pl.col("close"))
.otherwise(pl.col(col))
.alias(col)
)
return result
@staticmethod
def _build_quote_extra(records: list[dict]) -> pl.DataFrame:
"""构建 API 直接提供的补充字段 (不写 daily, 只传给 enriched 计算)。
包含: prev_close, change_pct, change_amount, amplitude, turnover_rate。
"""
if not records:
return pl.DataFrame()
df = pl.DataFrame(records)
keep = [c for c in [
"symbol", "prev_close", "change_pct", "change_amount",
"amplitude", "turnover_rate",
] if c in df.columns]
if not keep or "symbol" not in keep:
return pl.DataFrame()
return df.select(keep)
@staticmethod
def _build_index_quotes(records: list[dict]) -> pl.DataFrame:
"""构建指数实时行情缓存,不落股票 parquet。
注意: API 返回的 change_pct/amplitude 是小数 (0.0366 = 3.66%),
统一转成百分比输出, 与 _fallback_index_quotes_from_daily 口径一致
(前端指数侧不×100, 直接 toFixed(2)% 展示)。
"""
if not records:
return pl.DataFrame()
df = pl.DataFrame(records)
keep = [c for c in [
"symbol", "name", "last_price", "prev_close", "open", "high", "low",
"volume", "amount", "change_pct", "change_amount", "amplitude", "timestamp", "session",
] if c in df.columns]
if not keep or "symbol" not in keep:
return pl.DataFrame()
df = df.select(keep)
# change_pct / amplitude: 小数 → 百分比 (统一指数展示口径)
for col in ("change_pct", "amplitude"):
if col in df.columns:
df = df.with_columns((pl.col(col).cast(pl.Float64) * 100).alias(col))
if "last_price" in df.columns and "close" not in df.columns:
df = df.with_columns(pl.col("last_price").alias("close"))
return df
@staticmethod
def _is_trading_hours() -> bool:
now = datetime.now()
t = now.time()
morning = dt_time(9, 15) <= t <= dt_time(11, 35)
afternoon = dt_time(12, 55) <= t <= dt_time(15, 5)
return now.weekday() < 5 and (morning or afternoon)
@staticmethod
def _save_enabled(enabled: bool) -> None:
from app.services import preferences
preferences.save({"realtime_quotes_enabled": enabled})
# ================================================================
# 策略监控
# ================================================================
def _evaluate_monitors(self, daily_df: pl.DataFrame, quote_extra: pl.DataFrame | None) -> None:
"""行情更新后评估统一监控规则引擎,并刷新策略结果缓存。"""
try:
# 获取 enriched 数据 (刚算好的)
enriched_today, enriched_date = self.get_enriched_today()
if enriched_today.is_empty():
return
all_alerts: list[dict] = []
rule_events: list[dict] = []
engine = None
# 通用监控规则评估 (统一引擎: signal/price/market/strategy)
if self._app_state:
engine = getattr(self._app_state, "monitor_engine", None)
if engine and engine.rule_count > 0:
# 预构建 symbol → name 映射 (enriched 已 drop name 列, 引擎触发时回填用)
try:
inst_df = self._app_state.repo.get_instruments()
if not inst_df.is_empty() and "symbol" in inst_df.columns and "name" in inst_df.columns:
engine.set_name_map({
row["symbol"]: row["name"]
for row in inst_df.select(["symbol", "name"]).iter_rows(named=True)
if row.get("name")
})
except Exception as e: # noqa: BLE001
logger.debug("name_map 构建失败 (不影响监控): %s", e)
# 连板梯队封单监控: 有 ladder 规则时, 从 depth_service 注入封单量到 enriched
eval_df = enriched_today
if engine.has_rule_type("ladder"):
eval_df = self._inject_sealed_vol(enriched_today, enriched_date)
rule_events = engine.evaluate(eval_df)
if rule_events:
# 落盘到 alerts.jsonl
try:
from app.services import alert_store
alert_store.append_many(
self._app_state.repo.store.data_dir, rule_events,
)
except Exception as e: # noqa: BLE001
logger.warning("告警落盘失败: %s", e)
# 转为 SSE 推送格式 (兼容旧 alert schema)
for ev in rule_events:
all_alerts.append({
"source": ev["source"],
"type": ev["type"],
"rule_id": ev.get("rule_id"),
"strategy_id": ev.get("rule_id") if ev["source"] == "strategy" else None,
"symbol": ev["symbol"],
"name": ev["name"],
"message": ev["message"],
"price": ev["price"],
"change_pct": ev["change_pct"],
"signals": ev["signals"],
"severity": ev.get("severity", "info"),
"conditions": ev.get("conditions") or [],
"logic": ev.get("logic") or "and",
})
# 策略页实时回显: 不写文件 (实时行情每轮更新 enriched, 写文件会被 read_cache
# 的 mtime 校验判过期, 反复读不到)。监控引擎本轮已算出的结果存在内存
# (latest_strategy_results), 由 /api/screener/cached 端点直接叠加读取。
# 推入待推送队列 + 通知 SSE (含背压保护)
if all_alerts:
with self._lock:
self._pending_alerts.extend(all_alerts)
# 背压: 超出上限丢弃最旧
if len(self._pending_alerts) > self._max_pending_alerts:
overflow = len(self._pending_alerts) - self._max_pending_alerts
self._pending_alerts = self._pending_alerts[overflow:]
self._alert_event.set()
logger.info("监控评估完成: %d 条通知", len(all_alerts))
# 系统通知 (可选通道, 由 preferences 开关控制)。
# cooldown 去重已在 MonitorRuleEngine 做过, 这里只负责转发。
self._maybe_send_system_notifications(all_alerts)
# Webhook 推送 (飞书等外部 IM, 由规则 webhook_enabled 开关控制)。
# 紧随系统通知, 同样静默降级不阻断主流程。
if rule_events:
self._maybe_send_webhook(rule_events, engine)
except Exception as e: # noqa: BLE001
logger.warning("监控评估失败: %s", e)
def _inject_sealed_vol(self, enriched_today: pl.DataFrame, enriched_date) -> pl.DataFrame:
"""从 depth_service 取封单量, 作为临时列 _sealed_vol 注入 enriched 副本。
涨停封单(买一量) + 跌停封单(卖一量)合并, 供 ladder 规则评估。
depth 未就绪时返回原 df (不注入, ladder 规则安全降级不触发)。
"""
try:
depth_svc = getattr(self._app_state, "depth_service", None)
if not depth_svc:
return enriched_today
# enriched_date 可能是 date 或字符串, 统一为 date
from datetime import date as date_cls
target_date = enriched_date if isinstance(enriched_date, date_cls) else date_cls.fromisoformat(str(enriched_date))
# 取涨停 + 跌停封单, 合并 {symbol: vol}
up_map = depth_svc.get_sealed_map(target_date, is_down=False)
down_map = depth_svc.get_sealed_map(target_date, is_down=True)
sealed: dict[str, int] = {}
for m in (up_map, down_map):
for sym, info in m.items():
vol = (info or {}).get("vol")
if vol and vol > 0:
sealed[sym] = vol # 后者覆盖前者 (同 symbol 不可能在涨跌停都封单)
if not sealed:
return enriched_today
# 构造 (symbol, _sealed_vol) DataFrame, join 到 enriched 副本
sealed_df = pl.DataFrame({
"symbol": list(sealed.keys()),
"_sealed_vol": list(sealed.values()),
})
# 若已有残留列先移除 (避免重复 join 报错)
df = enriched_today.drop("_sealed_vol") if "_sealed_vol" in enriched_today.columns else enriched_today
return df.join(sealed_df, on="symbol", how="left")
except Exception as e: # noqa: BLE001
logger.debug("封单注入失败 (ladder 规则将不触发): %s", e)
return enriched_today
def _maybe_send_webhook(self, rule_events: list[dict], engine) -> None:
"""把告警通过 Webhook 推送到外部 IM (由规则 webhook_enabled 开关控制)。
- 全局飞书 URL 未配置: 直接返回
- 仅推送 webhook_enabled=True 的规则触发的告警
- 失败静默, 不阻断主流程
- 去重: 复用 MonitorRuleEngine 的 cooldown, 此处不重复去重
注意: 用 rule_events (含 rule_id) 而非重建后的 all_alerts,
以便反查引擎规则判断是否启用推送。
"""
try:
from app.services import preferences
from app.services import webhook_adapter
url = preferences.get_feishu_webhook_url()
if not url:
return
secret = preferences.get_feishu_webhook_secret()
# 反查规则, 过滤出启用推送的事件
source_labels = {
"strategy": "策略", "signal": "信号",
"price": "价格", "market": "异动",
}
rules = engine.rules if engine is not None else {}
pushed = 0
for ev in rule_events:
rule = rules.get(ev.get("rule_id"))
if not rule or not rule.get("webhook_enabled"):
continue
source = ev.get("source", "")
source_label = source_labels.get(source, source or "通知")
symbol = ev.get("symbol") or ""
name = ev.get("name") or ""
message = ev.get("message") or ""
title = f"TickFlow · {source_label}"
body = f"{symbol} {name} {message}".strip() if symbol else (message or name)
if webhook_adapter.send_feishu(url, title, body, secret):
pushed += 1
if pushed:
logger.info("飞书 Webhook 推送: %d", pushed)
except Exception as e: # noqa: BLE001
logger.debug("Webhook 推送异常 (不影响告警主流程): %s", e)
def _maybe_send_system_notifications(self, all_alerts: list[dict]) -> None:
"""把告警转发到操作系统通知中心 (由 preferences 开关控制)。
- 开关关闭: 直接返回
- 开关开启: 逐条发系统通知; 失败静默, 不阻断主流程
- 去重: 复用 MonitorRuleEngine 的 cooldown, 此处不重复去重
- 批量策略事件 (symbol="") 聚合为一条通知, 避免刷屏
"""
try:
from app.services import preferences
from app.services import notify_adapter
if not preferences.get_system_notify_enabled():
return
for ev in all_alerts:
# 通知标题: 用 source 分类 (策略/信号/价格/异动)
source = ev.get("source", "")
source_label = {
"strategy": "策略", "signal": "信号",
"price": "价格", "market": "异动",
}.get(source, source or "通知")
name = ev.get("name") or ""
symbol = ev.get("symbol") or ""
message = ev.get("message") or ""
# 正文: 优先用现成 message, 拼上 symbol/name 让用户一眼定位
if symbol:
body = f"{symbol} {name} {message}".strip()
else:
body = message or name
title = f"TickFlow · {source_label}"
notify_adapter.notify(title, body)
except Exception as e: # noqa: BLE001
logger.debug("系统通知发送异常 (不影响告警主流程): %s", e)
@staticmethod
def _get_strategy_monitor():
"""获取 StrategyMonitorService — 不再使用, 改用 _app_state 注入。"""
return None
# ================================================================
# enriched 增量计算
# ================================================================
def _flush_live_enriched(self, daily_df: pl.DataFrame, quote_extra: pl.DataFrame = None, asset_type: str = "stock", merge: bool = False) -> None:
"""增量计算今天的 enriched: 用昨天的递推状态 + 今天 OHLCV → 只算今天 5500 行。
quote_extra: API 直接提供的补充字段 (prev_close, change_pct 等),
不写 daily, 直接传给 compute_enriched_today 避免重复计算。
"""
try:
today = date.today()
t0 = time.perf_counter()
# ---- 尝试增量路径 ----
live_agg = self._repo.get_live_agg() if asset_type == "stock" else pl.DataFrame()
prev_enriched, prev_date = (
self._repo.get_enriched_latest()
if asset_type == "stock"
else self._repo.get_enriched_latest_asset(asset_type)
)
use_incremental = (
asset_type == "stock"
and not live_agg.is_empty()
and not prev_enriched.is_empty()
and prev_date is not None
)
if use_incremental:
from app.indicators.pipeline import compute_enriched_today
instruments = self._repo.get_instruments()
# 将 API 直接提供的补充字段 JOIN 到 daily_df
today_ohlcv = daily_df
if quote_extra is not None and not quote_extra.is_empty():
today_ohlcv = daily_df.join(quote_extra, on="symbol", how="left")
enriched_today = compute_enriched_today(
live_agg=live_agg,
prev_enriched=prev_enriched,
today_ohlcv=today_ohlcv,
instruments=instruments,
)
if enriched_today.is_empty():
logger.warning("增量计算结果为空, 回退到全量计算")
use_incremental = False
# ---- 全量回退路径 ----
if not use_incremental:
from datetime import timedelta
from app.indicators.pipeline import compute_enriched
logger.info("enriched 全量计算 (live_agg=%s, 上次日期=%s)",
"ok" if not live_agg.is_empty() else "", prev_date)
cutoff = today - timedelta(days=90)
table = "kline_etf_daily" if asset_type == "etf" else "kline_daily"
daily_glob = str(self._repo.store.data_dir / table / "**" / "*.parquet")
ohlcv_cols = ["symbol", "date", "open", "high", "low", "close", "volume", "amount"]
hist_df = (
pl.scan_parquet(daily_glob)
.filter(pl.col("date") >= cutoff)
.sort(["symbol", "date"])
.collect()
)
if hist_df.is_empty():
return
hist_cols = [c for c in ohlcv_cols if c in hist_df.columns]
hist_df = hist_df.select(hist_cols).filter(pl.col("date") != today)
daily_ohlcv = daily_df.select([c for c in ohlcv_cols if c in daily_df.columns])
full_df = pl.concat([hist_df, daily_ohlcv], how="diagonal_relaxed")
full_df = full_df.sort(["symbol", "date"])
factor_dir = "adj_factor_etf" if asset_type == "etf" else "adj_factor"
factor_path = self._repo.store.data_dir / factor_dir / "all.parquet"
factors = pl.DataFrame()
if factor_path.exists():
try:
factors = pl.read_parquet(factor_path)
except Exception:
pass
instruments = self._repo.get_instruments() if asset_type == "stock" else None
enriched_full = compute_enriched(full_df, factors=factors, instruments=instruments)
enriched_today = enriched_full.filter(pl.col("date") == today)
if enriched_today.is_empty():
return
# ---- 写盘 + 更新缓存 ----
if merge:
self._repo.merge_live_enriched_asset(asset_type, enriched_today)
else:
self._repo.flush_live_enriched_asset(asset_type, enriched_today)
elapsed = time.perf_counter() - t0
mode_label = "增量" if use_incremental else "全量"
logger.info("enriched %s: %d 只, %s, 耗时 %.0fms",
mode_label, len(enriched_today), today, elapsed * 1000)
except Exception as e: # noqa: BLE001
logger.warning("enriched 计算失败: %s", e)
+214
View File
@@ -0,0 +1,214 @@
"""概念涨幅轮动矩阵 service。
输出「每列(日期)各自把所有概念按当天涨幅从高到低排序」的矩阵,供前端
「概念分析 → 涨幅RPS轮动」对话框渲染。
数据来源全部复用现有资产, 不引入新数据源:
- 个股历史涨跌幅: repo.get_enriched_range(..., columns=["symbol","date","change_pct"])
命中启动时构建的 _enriched_history_cache (0ms, 含 change_pct 小数列)
- 概念成分股映射: 复用 market_overview_builder 的 _dimension_field / _read_ext_rows /
_symbol_keys / _dimension_values, 与看板/复盘的概念聚合口径完全一致
性能: 387 概念 × 30 天的 group_by + sort 是 polars 内存操作, 实测 <50ms;
另加进程级结果缓存 (_CACHE_TTL=120s), 重复请求 <1ms。
"""
from __future__ import annotations
import logging
import time
from datetime import date, timedelta
import polars as pl
from app.services.market_overview_builder import (
_dimension_field,
_dimension_values,
_read_ext_rows,
_symbol_keys,
)
from app.services.ext_data import ExtConfigStore
logger = logging.getLogger(__name__)
# 进程级结果缓存 (照搬 overview.py:18 的模式, TTL 拉长到 120s —— 轮动矩阵
# 不像看板那样需要近实时, 盘后数据稳定, 缓存久一点无妨)
_CACHE_TTL = 120.0
_cache: dict[str, dict] = {}
_cache_ts: dict[str, float] = {}
def invalidate_cache() -> None:
"""清空轮动矩阵结果缓存(数据管道完成后调用, 避免返回旧数据)。"""
_cache.clear()
_cache_ts.clear()
def _latest_enriched_date(repo) -> date | None:
"""取 enriched 缓存里的最新交易日(矩阵的右端=最新日期)。"""
cache = repo._enriched_history_cache # noqa: SLF001 —— 缓存字段无公开 getter
if cache is None or cache.is_empty() or "date" not in cache.columns:
return None
return cache["date"].max()
def _load_concept_map_df(repo) -> tuple[pl.DataFrame, int]:
"""构建并缓存 {symbol_upper → 概念} 的已展开 polars 映射表。
复用 market_overview_builder 的概念识别 + 成分股读取逻辑(_dimension_field /
_read_ext_rows / _symbol_keys / _dimension_values), 但要的是「反向映射」
(symbol → 概念), 且直接产出 polars DataFrame 供 join 使用。
返回 (map_df, concept_count):
- map_df: 两列 (_sym_up: 大写 symbol, concept: 概念名), 已 explode, 一个
symbol 属多概念时有多行。无概念数据时返回空 DataFrame。
- concept_count: 去重概念总数。
缓存: 概念成分股是 snapshot, 进程内不变, 缓存 600s。
直接缓存 DataFrame 而非 Python dict —— 后续 join 时省掉每次 ~1s 的 dict→DataFrame
重建开销(这是结果缓存失效后重算的主要瓶颈)。
"""
global _concept_map_cache, _concept_map_count, _concept_map_ts
now = time.time()
if _concept_map_cache is not None and (now - _concept_map_ts) < 600:
return _concept_map_cache, _concept_map_count
data_dir = repo.store.data_dir
store = ExtConfigStore(data_dir)
# 先收集成扁平的 (sym, concept) 行, 再一次性构造 DataFrame(比 list 列快得多)
pairs: list[tuple[str, str]] = []
concepts_seen: set[str] = set()
for config in store.load_all():
field = _dimension_field(config, "concept")
if not field:
continue
for ext_row in _read_ext_rows(data_dir, config, field):
concepts = _dimension_values(ext_row.get(field))
if not concepts:
continue
keys = _symbol_keys(ext_row, config)
for key in keys:
for c in concepts:
pairs.append((key, c))
concepts_seen.add(c)
if pairs:
# 去重: 同一 (symbol, concept) 对会因多 key 形式(SZ/000001)和
# 多 config 重复出现, 去重后从 ~48万 行降到 ~14万, join 快 3x+
_concept_map_cache = pl.DataFrame(
{"_sym_up": [p[0] for p in pairs], "concept": [p[1] for p in pairs]},
schema={"_sym_up": pl.Utf8, "concept": pl.Utf8},
).unique()
_concept_map_count = len(concepts_seen)
else:
_concept_map_cache = pl.DataFrame(
schema={"_sym_up": pl.Utf8, "concept": pl.Utf8}
)
_concept_map_count = 0
_concept_map_ts = now
return _concept_map_cache, _concept_map_count
_concept_map_cache: pl.DataFrame | None = None
_concept_map_count: int = 0
_concept_map_ts: float = 0.0
def build_rps_rotation(repo, days: int = 12) -> dict:
"""构建概念涨幅轮动矩阵。
Args:
repo: KlineRepository(含 _enriched_history_cache 内存历史)。
days: 取最近 N 个交易日, 范围 [7, 30], 默认 12。
Returns:
{
"dates": ["2026-06-30", ...], # 最新在最前, 长度 ≤ days
"columns": {"2026-06-30": [[概念, 涨幅], ...], ...}, # 每列各自排序(高→低)
"concept_count": 387, # 去重概念总数(0 表示无概念数据)
}
涨幅是小数(0.0522 = +5.22%)。无数据时返回空 columns。
"""
days = max(7, min(30, days))
# 结果缓存: 同 days(→ 同 start/end)的请求在 TTL 内直接返回
latest = _latest_enriched_date(repo)
if latest is None:
return {"dates": [], "columns": {}, "concept_count": 0}
cache_key = latest.isoformat()
now = time.time()
cached = _cache.get(cache_key)
if cached and (now - _cache_ts.get(cache_key, 0)) < _CACHE_TTL:
# 缓存的是所有日期, 按需要的 days 截取(避免不同 days 各存一份)
return _slice_cached(cached, days)
# 1. 概念映射(symbol → 概念), 已缓存为 polars DataFrame
map_df, concept_count = _load_concept_map_df(repo)
if map_df.is_empty():
logger.info("rps_rotation: no concept data (ext_gn_ths not fetched yet)")
return {"dates": [], "columns": {}, "concept_count": 0}
# 2. 取最近 N 交易日的个股 change_pct(命中内存缓存)
start = latest - timedelta(days=days * 2 + 10) # 日历天 ≈ 2/3 交易日, 多取余量
df = repo.get_enriched_range(
start, latest, columns=["symbol", "date", "change_pct"]
)
if df is None or df.is_empty():
return {"dates": [], "columns": {}, "concept_count": 0}
# 3. 把个股 symbol 映射到概念, 一只股票拆成多行(每个概念一行)
# symbol 大写匹配(map_df 的 _sym_up 已大写)
df = df.with_columns(pl.col("symbol").str.to_uppercase().alias("_sym_up"))
joined = df.join(map_df, on="_sym_up", how="inner").drop("_sym_up")
if joined.is_empty():
return {"dates": [], "columns": {}, "concept_count": 0}
# 4. 按 (date, concept) 聚合 avg change_pct —— 与 _dimension_rank:288 的简单平均口径一致
agg = joined.group_by(["date", "concept"]).agg(
pl.col("change_pct").mean().alias("avg_pct")
)
# 去掉 NaN/Null(停牌等无行情的概念日)
agg = agg.filter(pl.col("avg_pct").is_not_null() & pl.col("avg_pct").is_not_nan())
# 5. 每个日期内按 avg_pct 降序排, 再 group_by 把每组的 (concept, avg_pct)
# 收集成并行 list —— 一次 polars 操作拿到全部列, 避免 partition_by 的 tuple key 歧义
agg = agg.sort(["date", "avg_pct"], descending=[False, True])
grouped = agg.group_by("date", maintain_order=True).agg(
pl.col("concept"), pl.col("avg_pct")
)
# 最新日期排最前
grouped = grouped.sort("date", descending=True)
columns: dict[str, list[list]] = {}
all_dates_sorted: list[str] = []
for row in grouped.iter_rows(named=True):
d_str = str(row["date"])
all_dates_sorted.append(d_str)
columns[d_str] = list(zip(row["concept"], row["avg_pct"]))
full = {
"dates": [str(d) for d in all_dates_sorted],
"columns": columns,
"concept_count": concept_count,
}
# 写缓存(存全量, 按需 slice)
_cache[cache_key] = full
_cache_ts[cache_key] = now
return _slice_cached(full, days)
def _slice_cached(full: dict, days: int) -> dict:
"""从全量缓存截取最近 N 天(days)。"""
dates_all = full["dates"]
if len(dates_all) <= days:
return full
keep_dates = dates_all[:days]
return {
"dates": keep_dates,
"columns": {d: full["columns"][d] for d in keep_dates},
"concept_count": full["concept_count"],
}
+593
View File
@@ -0,0 +1,593 @@
"""Screener 服务(§6.3)。
性能优化:
- enriched parquet 仅存 14 列基础数据, 指标和信号即时计算
- preset 策略: 从内存缓存或即时计算获取完整指标, ~10-50ms
- custom SQL: DuckDB (用户传 SQL WHERE 字符串), ~10-50ms
"""
from __future__ import annotations
import logging
import time
from dataclasses import dataclass, field
from datetime import date, timedelta
import polars as pl
from app.tickflow.repository import KlineRepository
logger = logging.getLogger(__name__)
# ── 进程级历史数据缓存 (避免 run_all 每次重新扫描 parquet + 计算指标) ──
_history_cache: dict[tuple[date, int], tuple[float, pl.DataFrame]] = {}
_HISTORY_CACHE_TTL = 120.0 # 秒
# 内置预设策略 — Polars 表达式方式
PRESET_STRATEGIES: dict[str, dict] = {
"trend_breakout": {
"name": "趋势突破",
"description": "MA60 上方 + 60 日新高 + 量能 ≥ 2 倍均量",
"filter": (
(pl.col("close") > pl.col("ma60"))
& pl.col("signal_n_day_high").fill_null(False)
& (pl.col("vol_ratio_5d") >= 2.0)
),
"order_by": "momentum_60d",
"descending": True,
"limit": 100,
},
"ma_golden_cross": {
"name": "MA 金叉",
"description": "MA5 上穿 MA20 当日触发,量能配合",
"filter": (
pl.col("signal_ma_golden_5_20").fill_null(False)
& (pl.col("vol_ratio_5d") >= 1.2)
& (pl.col("close") > pl.col("ma60"))
),
"order_by": "momentum_20d",
"descending": True,
"limit": 100,
},
"macd_golden": {
"name": "MACD 金叉放量",
"description": "MACD 金叉当日 + 量能放大",
"filter": (
pl.col("signal_macd_golden").fill_null(False)
& (pl.col("vol_ratio_5d") >= 1.5)
),
"order_by": "momentum_60d",
"descending": True,
"limit": 100,
},
"volume_price_surge": {
"name": "量价齐升",
"description": "突破 MA20 + 放量 + 收阳",
"filter": (
pl.col("signal_ma20_breakout").fill_null(False)
& (pl.col("vol_ratio_5d") >= 2.0)
& (pl.col("close") > pl.col("open"))
),
"order_by": "vol_ratio_5d",
"descending": True,
"limit": 100,
},
"low_volatility_leader": {
"name": "低波动龙头",
"description": "20 日动量为正 + 年化波动 < 30% + MA20 上方",
"filter": (
(pl.col("momentum_20d") > 0)
& (pl.col("annual_vol_20d") < 0.30)
& (pl.col("close") > pl.col("ma20"))
),
"order_by": "momentum_60d",
"descending": True,
"limit": 100,
},
"broken_board_recovery": {
"name": "断板反包",
"description": "连板 ≥2 后断板 1-2 天,出现放量反包信号",
"filter": (
pl.col("signal_limit_up").fill_null(False)
& (pl.col("vol_ratio_5d") >= 1.5)
& (pl.col("change_pct") > 0.03)
),
"order_by": "change_pct",
"descending": True,
"limit": 100,
},
"oversold_bounce": {
"name": "超跌反弹",
"description": "RSI14 < 30 超卖区 + 当日收阳 + 放量,抄底信号",
"filter": (
(pl.col("rsi_14") < 30)
& (pl.col("close") > pl.col("open"))
& (pl.col("vol_ratio_5d") >= 1.2)
),
"order_by": "rsi_14",
"descending": False,
"limit": 100,
},
"boll_breakout": {
"name": "布林突破",
"description": "突破布林上轨 + 放量,强势加速信号",
"filter": (
pl.col("signal_boll_breakout_upper").fill_null(False)
& (pl.col("vol_ratio_5d") >= 1.5)
),
"order_by": "vol_ratio_5d",
"descending": True,
"limit": 100,
},
"bullish_alignment": {
"name": "均线多头",
"description": "MA5 > MA10 > MA20 > MA60 多头排列 + 短期动量为正",
"filter": (
(pl.col("ma5") > pl.col("ma10"))
& (pl.col("ma10") > pl.col("ma20"))
& (pl.col("ma20") > pl.col("ma60"))
& (pl.col("momentum_20d") > 0)
),
"order_by": "momentum_60d",
"descending": True,
"limit": 100,
},
"consecutive_limit_ups": {
"name": "连板股",
"description": "当日涨停且连续涨停 ≥ 2 天,强势追涨",
"filter": (
pl.col("signal_limit_up").fill_null(False)
& (pl.col("consecutive_limit_ups") >= 2)
),
"order_by": "consecutive_limit_ups",
"descending": True,
"limit": 100,
},
"pullback_to_support": {
"name": "缩量回踩",
"description": "回踩 MA20 附近 + 缩量 + 中期趋势向上",
"filter": (
(pl.col("close") > pl.col("ma20") * 0.98)
& (pl.col("close") < pl.col("ma20") * 1.02)
& (pl.col("vol_ratio_5d") < 0.8)
& (pl.col("close") > pl.col("ma60"))
& (pl.col("momentum_20d") > 0)
),
"order_by": "momentum_60d",
"descending": True,
"limit": 100,
},
"n_day_low_reversal": {
"name": "新低反转",
"description": "触及 60 日新低后当日收阳放量,反转信号",
"filter": (
pl.col("signal_n_day_low").fill_null(False)
& (pl.col("close") > pl.col("open"))
& (pl.col("vol_ratio_5d") >= 1.5)
),
"order_by": "change_pct",
"descending": True,
"limit": 100,
},
}
@dataclass
class ScreenerResult:
as_of: date
strategy: str | None
rows: list[dict] = field(default_factory=list)
total: int = 0
elapsed_ms: float = 0.0
class ScreenerService:
def __init__(self, repo: KlineRepository) -> None:
self.repo = repo
@staticmethod
def clear_history_cache() -> None:
"""清空进程级 _history_cache (TTL 缓存)。
清除数据后调用, 避免内存里的旧历史窗口残留导致策略/看板仍命中旧数据。
"""
_history_cache.clear()
def _load_enriched_for_date(self, target_date: date) -> pl.DataFrame:
"""从 enriched parquet 读取指定日期的基础数据并即时计算完整指标+信号。
enriched parquet 仅存 14 列。读取后需要即时计算 ma/ema/macd/kdj/rsi/boll/momentum/signal 等列。
对于最新日, 优先使用内存缓存 (已包含完整指标)。
"""
# 优先使用 repo 最新日缓存
cache, cache_date = self.repo.get_enriched_latest()
if cache is not None and not cache.is_empty() and cache_date == target_date:
df = cache
# JOIN instruments
df_i = self.repo.get_instruments()
if not df_i.is_empty():
inst_cols = [c for c in ["symbol", "name", "total_shares", "float_shares"] if c in df_i.columns]
if "name" not in df.columns:
df = df.join(df_i.select(inst_cols), on="symbol", how="left")
return df
# 尝试从 repo 级预计算历史缓存中提取目标日期
cached_hist = self.repo.get_enriched_history(target_date, 1)
if cached_hist is not None and not cached_hist.is_empty() and "date" in cached_hist.columns:
df = cached_hist.filter(pl.col("date") == target_date)
if not df.is_empty():
logger.debug("_load_enriched_for_date: repo history cache for %s", target_date)
# JOIN instruments
df_i = self.repo.get_instruments()
if not df_i.is_empty():
inst_cols = [c for c in ["symbol", "name", "total_shares", "float_shares"] if c in df_i.columns]
if "name" not in df.columns:
df = df.join(df_i.select(inst_cols), on="symbol", how="left")
return df
# 历史日期: 从 parquet 读取 14 列, 即时计算指标 (慢路径)
enriched_dir = self.repo.store.data_dir / "kline_daily_enriched"
ds = target_date.isoformat()
target_parquet = enriched_dir / f"date={ds}" / "part.parquet"
if not target_parquet.exists():
return pl.DataFrame()
try:
df = pl.read_parquet(target_parquet)
except Exception as e: # noqa: BLE001
logger.warning("load_enriched_for_date failed: %s", e)
return pl.DataFrame()
if df.is_empty():
return df
# 即时计算指标: 需要加载历史窗口作 warmup
df_full = self._compute_enriched_full(df, target_date)
return df_full
def _compute_enriched_full(self, df_target: pl.DataFrame, target_date: date) -> pl.DataFrame:
"""从 14 列基础数据即时计算完整 enriched (含全部指标和信号)。
读取历史数据作为指标计算的 warmup, 计算完成后只返回目标日期的行。
"""
from app.indicators.pipeline import compute_indicators, compute_signals, compute_limit_signals
# 加载 warmup 历史 (目标日期前 ~120 天)
enriched_dir = self.repo.store.data_dir / "kline_daily_enriched"
start = target_date - timedelta(days=150)
read_cols = ["symbol", "date", "open", "high", "low", "close", "volume",
"amount", "raw_close", "raw_high", "raw_low"]
try:
lf = (
pl.scan_parquet(str(enriched_dir / "**" / "*.parquet"))
.filter(
(pl.col("date") >= start)
& (pl.col("date") <= target_date)
)
.sort(["symbol", "date"])
)
available = [c for c in read_cols if c in lf.schema]
df_hist = lf.select(available).collect()
except Exception as e: # noqa: BLE001
logger.warning("warmup history load failed: %s", e)
df_hist = df_target
if df_hist.is_empty():
df_hist = df_target
# 计算指标
df_full = compute_indicators(df_hist)
df_full = compute_signals(df_full)
# 计算涨跌停信号 (需要 instruments)
instruments = self.repo.get_instruments()
if instruments is not None and not instruments.is_empty():
df_full = compute_limit_signals(df_full, instruments)
# 只保留目标日期
df_result = df_full.filter(pl.col("date") == target_date)
# JOIN instruments (name, total_shares, float_shares)
if not instruments.is_empty():
inst_cols = [c for c in ["symbol", "name", "total_shares", "float_shares"] if c in instruments.columns]
if "name" not in df_result.columns:
df_result = df_result.join(instruments.select(inst_cols), on="symbol", how="left")
return df_result
def _load_enriched_history(self, target_date: date, lookback_days: int) -> pl.DataFrame:
"""读取目标日期之前的基础行情数据, 供历史窗口策略使用。
优先从 repo 内存缓存获取 (启动时已预计算), 命中时 0ms。
缓存 miss 时走 scan_parquet + compute_indicators 慢路径。
"""
# 优先级 1: repo 级预计算缓存 (启动时 _refresh_enriched 已计算完整历史)
t0 = time.perf_counter()
cached = self.repo.get_enriched_history(target_date, lookback_days)
if cached is not None and not cached.is_empty():
# JOIN instruments (repo 缓存不含 name 等列)
instruments = self.repo.get_instruments()
if instruments is not None and not instruments.is_empty() and "name" not in cached.columns:
inst_cols = [c for c in ["symbol", "name", "total_shares", "float_shares"]
if c in instruments.columns]
cached = cached.join(instruments.select(inst_cols), on="symbol", how="left")
elapsed = (time.perf_counter() - t0) * 1000
logger.info("_load_enriched_history(%s, %d): repo cache hit, %.1fms, %d rows",
target_date, lookback_days, elapsed, len(cached))
return cached
# 优先级 2: 进程级 history_cache (之前的 TTL 缓存)
cache_key = (target_date, lookback_days)
now = time.monotonic()
ttl_cached = _history_cache.get(cache_key)
if ttl_cached is not None:
ts, cached_df = ttl_cached
if now - ts < _HISTORY_CACHE_TTL:
logger.debug("history TTL cache hit: %s lookback=%d", target_date, lookback_days)
return cached_df
del _history_cache[cache_key]
# 优先级 3: scan_parquet + compute_indicators (慢路径, ~5s)
logger.warning("_load_enriched_history cache miss, computing indicators (%s, %d)...",
target_date, lookback_days)
from app.indicators.pipeline import compute_indicators, compute_signals, compute_limit_signals
warmup = 60
start = target_date - timedelta(days=min((lookback_days + warmup) * 2, 180))
enriched_dir = self.repo.store.data_dir / "kline_daily_enriched"
read_cols = ["symbol", "date", "open", "high", "low", "close", "volume",
"amount", "raw_close", "raw_high", "raw_low"]
try:
lf = (
pl.scan_parquet(str(enriched_dir / "**" / "*.parquet"))
.filter((pl.col("date") >= start) & (pl.col("date") <= target_date))
.sort(["symbol", "date"])
)
available = [c for c in read_cols if c in lf.collect_schema().names()]
df_hist = lf.select(available).collect()
except Exception as e: # noqa: BLE001
logger.warning("load_enriched_history failed: %s", e)
return pl.DataFrame()
if df_hist.is_empty():
return pl.DataFrame()
df_full = compute_indicators(df_hist)
df_full = compute_signals(df_full)
instruments = self.repo.get_instruments()
if instruments is not None and not instruments.is_empty():
df_full = compute_limit_signals(df_full, instruments)
if instruments is not None and not instruments.is_empty():
inst_cols = [c for c in ["symbol", "name", "total_shares", "float_shares"] if c in instruments.columns]
if "name" not in df_full.columns:
df_full = df_full.join(instruments.select(inst_cols), on="symbol", how="left")
# 裁剪掉 warmup 部分, 只保留 lookback 范围 (减少 group_by 开销)
lookback_start = target_date - timedelta(days=lookback_days)
if "date" in df_full.columns:
df_full = df_full.filter(pl.col("date") >= lookback_start)
df_full = df_full.sort(["symbol", "date"])
elapsed = (time.perf_counter() - t0) * 1000
logger.info("_load_enriched_history(%s, %d): computed in %.1fms, %d rows",
target_date, lookback_days, elapsed, len(df_full))
_history_cache[cache_key] = (now, df_full)
if len(_history_cache) > 10:
expired = [k for k, (ts, _) in _history_cache.items() if now - ts > _HISTORY_CACHE_TTL]
for k in expired:
del _history_cache[k]
return df_full
def run(
self,
as_of: date,
conditions: list[str],
order_by: str | None = None,
limit: int = 30,
pool: list[str] | None = None,
) -> ScreenerResult:
"""自定义 SQL 条件选股。
先通过 Polars 即时计算完整指标, 再用 DuckDB 做 SQL WHERE 过滤。
kline_enriched DuckDB 视图只有 14 列, 不能直接用于指标过滤。
"""
t0 = time.perf_counter()
if not conditions:
return ScreenerResult(as_of=as_of, strategy=None)
# 从即时计算获取完整 enriched 数据
df = self._load_enriched_for_date(as_of)
if df.is_empty():
return ScreenerResult(as_of=as_of, strategy=None)
# Pool 过滤
if pool:
df = df.filter(pl.col("symbol").is_in(pool))
# 用 DuckDB 做 SQL 过滤 (注册临时视图)
try:
import duckdb
con = duckdb.connect(database=":memory:")
con.register("enriched", df.to_arrow())
where = " AND ".join(f"({c})" for c in conditions)
sql = f"SELECT * FROM enriched WHERE {where}"
if order_by:
sql += f" ORDER BY {order_by}"
if limit:
sql += f" LIMIT {limit}"
df_result = con.execute(sql).pl()
con.close()
except Exception as e: # noqa: BLE001
logger.warning("screener SQL query failed: %s", e)
df_result = pl.DataFrame()
rows = df_result.to_dicts() if not df_result.is_empty() else []
elapsed = (time.perf_counter() - t0) * 1000
return ScreenerResult(
as_of=as_of,
strategy=None,
rows=rows,
total=len(rows),
elapsed_ms=elapsed,
)
def run_preset(
self,
strategy_id: str,
as_of: date,
pool: list[str] | None = None,
precomputed: pl.DataFrame | None = None,
basic_filter: dict | None = None,
display_limit: int | None = None,
) -> ScreenerResult:
"""预设策略选股 — 从 enriched 读取预计算好的指标列后过滤。
- precomputed 不为空: 直接复用(run_all 场景)
- precomputed 为空: 从 enriched 读目标日期
- basic_filter: 用户保存的基础参数过滤(boards、价格等)
"""
t0 = time.perf_counter()
strat = PRESET_STRATEGIES.get(strategy_id)
if not strat:
raise ValueError(f"unknown strategy: {strategy_id}")
if precomputed is not None and not precomputed.is_empty():
df = precomputed
else:
df = self._load_enriched_for_date(as_of)
if df.is_empty():
return ScreenerResult(as_of=as_of, strategy=strategy_id)
# 应用用户基础参数过滤(boards、价格区间等)
if basic_filter and basic_filter.get("enabled", True):
df = self._apply_basic_filter(df, basic_filter)
# 应用策略过滤
df = df.filter(strat["filter"])
# 应用 pool
if pool:
df = df.filter(pl.col("symbol").is_in(pool))
# 排序 + 限制
order_col = strat["order_by"]
if order_col in df.columns:
df = df.sort(order_col, descending=strat.get("descending", True))
# display_limit: None=不限制, 0=全部, N=前N个
if display_limit == 0:
limit = None # 不限制
elif display_limit is not None:
limit = display_limit
else:
limit = None # 未配置时默认不限制
if limit is not None and limit > 0:
df = df.head(limit)
# 基于排序列生成 0-100 评分 (与 StrategyEngine 统一)
if order_col in df.columns and not df.is_empty():
col_vals = df[order_col].cast(pl.Float64)
col_min = col_vals.min()
col_max = col_vals.max()
col_range = col_max - col_min
if col_range and col_range > 0:
normalized = (col_vals - col_min) / col_range
else:
normalized = pl.Series("norm", [0.5] * len(df))
if not strat.get("descending", True):
normalized = 1.0 - normalized
df = df.with_columns((normalized * 100).alias("score"))
rows = df.to_dicts()
elapsed = (time.perf_counter() - t0) * 1000
# sanitize
for r in rows:
for k, v in list(r.items()):
if isinstance(v, float) and (v != v or abs(v) == float("inf")):
r[k] = None
return ScreenerResult(
as_of=as_of,
strategy=strategy_id,
rows=rows,
total=len(rows),
elapsed_ms=elapsed,
)
@staticmethod
def _apply_basic_filter(df: pl.DataFrame, bf: dict) -> pl.DataFrame:
"""应用用户基础参数过滤(boards、价格区间、市值等)"""
exprs: list[pl.Expr] = []
if bf.get("price_min") is not None:
exprs.append(pl.col("close") >= bf["price_min"])
if bf.get("price_max") is not None:
exprs.append(pl.col("close") <= bf["price_max"])
if bf.get("float_cap_min") is not None and "float_shares" in df.columns:
exprs.append(pl.col("close") * pl.col("float_shares") >= bf["float_cap_min"])
if bf.get("float_cap_max") is not None and "float_shares" in df.columns:
exprs.append(pl.col("close") * pl.col("float_shares") <= bf["float_cap_max"])
if bf.get("amount_min") is not None:
exprs.append(pl.col("amount") >= bf["amount_min"])
if bf.get("amount_max") is not None:
exprs.append(pl.col("amount") <= bf["amount_max"])
if bf.get("turnover_min") is not None and "turnover_rate" in df.columns:
exprs.append(pl.col("turnover_rate") >= bf["turnover_min"])
if bf.get("turnover_max") is not None and "turnover_rate" in df.columns:
exprs.append(pl.col("turnover_rate") <= bf["turnover_max"])
if bf.get("exclude_st") and "name" in df.columns:
exprs.append(~pl.col("name").str.contains("(?i)ST|\\*ST|退"))
# 板块过滤
boards = bf.get("boards")
if boards and isinstance(boards, list) and len(boards) > 0:
board_exprs: list[pl.Expr] = []
for b in boards:
if b == "沪主板":
board_exprs.append(pl.col("symbol").str.starts_with("60"))
elif b == "深主板":
board_exprs.append(
pl.col("symbol").str.starts_with("00")
| pl.col("symbol").str.starts_with("001")
)
elif b == "创业板":
board_exprs.append(
pl.col("symbol").str.starts_with("300")
| pl.col("symbol").str.starts_with("301")
)
elif b == "科创板":
board_exprs.append(pl.col("symbol").str.starts_with("688"))
elif b == "北交所":
board_exprs.append(pl.col("symbol").str.contains(r"\.BJ$"))
if board_exprs:
exprs.append(pl.any_horizontal(board_exprs))
if exprs:
return df.filter(pl.all_horizontal(exprs))
return df
def latest_date(self) -> date | None:
d = self.repo.enriched_latest_date()
if d:
return d
# 回退 DuckDB
try:
res = self.repo.execute_one(
"SELECT max(date) FROM kline_enriched",
)
if res and res[0]:
d = res[0]
return d if isinstance(d, date) else date.fromisoformat(str(d))
except Exception: # noqa: BLE001
return None
return None
@@ -0,0 +1,309 @@
"""AI 个股分析服务 — 技术面 / 基本面 / 财务面 / 消息面 四维综合分析。
职责:
组合一只股票的 K 线(含已算好的技术指标)+ 财务表 + 关键价位 →
拼装"实战派交易员"级系统提示词 → 流式调用 LLM → 逐 chunk 吐给前端。
与 financial_analyzer.py 的区别(刻意区分,非复用):
- 角色:A 股实战派交易员 / 技术分析师(非 CFA 财务分析师)
- 数据源:K 线 + 技术指标为主,财务表为辅(财务分析以财务表为主)
- 输出框架:技术面→基本面→财务面→消息面(四维),落点是买卖区间与操作建议
(财务分析的落点是财务质量评级)
不知道: HTTP、前端、配置持久化。
"""
from __future__ import annotations
import json
import logging
from pathlib import Path
from typing import AsyncIterator
import polars as pl
from app.indicators.levels import compute_levels, summarize_levels
from app.services.financial_sync import get_financial_df
logger = logging.getLogger(__name__)
# 注入最近多少根日 K(技术面分析样本)
_KLINE_WINDOW = 90
# 注入财务表的最近期数
_MAX_PERIODS = 4
# ================================================================
# 数据加载
# ================================================================
def _load_kline(repo, symbol: str) -> pl.DataFrame:
"""读取该标的最近 N 根日 K(已含技术指标 / 信号)。
repo: KlineRepository;走内存缓存,性能可控。
"""
from datetime import date, timedelta
end = date.today()
start = end - timedelta(days=_KLINE_WINDOW * 2) # 多取一些保证交易日够
df = repo.get_daily(symbol, start, end)
if df.is_empty():
return df
return df.tail(_KLINE_WINDOW)
def _clean_rows(df: pl.DataFrame, keep_cols: list[str]) -> list[dict]:
"""把 DataFrame 转成 JSON 安全的 dict 列表(只保留关键列 + 清洗 NaN/Inf + date→字符串)。
polars 的 date 列会变成 Python datetime.date,json.dumps 无法直接序列化,
必须转成 ISO 字符串,否则 json.dumps 抛 TypeError 让整个流静默失败。
"""
import datetime
import math
cols = [c for c in keep_cols if c in df.columns]
sub = df.select(cols)
rows = []
for rec in sub.to_dicts():
clean = {}
for k, v in rec.items():
if isinstance(v, float):
clean[k] = None if not math.isfinite(v) else round(v, 4)
elif isinstance(v, (datetime.date, datetime.datetime)):
clean[k] = v.isoformat()
else:
clean[k] = v
rows.append(clean)
return rows
def _load_financials(data_dir: Path, symbol: str) -> dict[str, list[dict]]:
"""读取该标的核心财务指标 + 利润表(只取最有信息量的两张表)。
财务面只需要关键指标(ROE / 增速 / 毛利率 等),不需要把 4 张表全塞进上下文
(那是 financial_analyzer 的职责)。这里取轻量,留给技术面更多 token。
"""
out: dict[str, list[dict]] = {}
for table in ("metrics", "income"):
df = get_financial_df(data_dir, table)
if df.is_empty():
out[table] = []
continue
df = df.filter(pl.col("symbol") == symbol)
if df.is_empty():
out[table] = []
continue
if "period_end" in df.columns:
df = df.sort("period_end", descending=True).head(2) # 只取最近 2 期
import math
rows = []
for rec in df.to_dicts():
clean = {}
for k, v in rec.items():
if k == "symbol":
continue
if isinstance(v, float):
clean[k] = None if not math.isfinite(v) else v
else:
clean[k] = v
rows.append(clean)
out[table] = rows
return out
# ================================================================
# 系统提示词 —— 实战派交易员四维框架(与财务分析明确区分)
# ================================================================
_SYSTEM_PROMPT = """你是一位拥有 15 年 A 股一线实战经验的资深交易员兼技术分析师,擅长从 K 线、量价、关键价位与基本面交叉验证中把握买卖时机。你的任务是:基于提供的个股数据,产出一份**实战、可直接指导交易决策**的综合分析报告。
## 输出规范
用 **Markdown** 格式输出,严格遵循以下结构。不要输出任何 JSON 或代码块,直接输出 Markdown 正文。
### 1. 🎯 一句话定调(1-2 句)
用一句话概括该股当前的**技术状态与交易属性**(如"高位放量滞涨,需警惕回调"/"底部筹码集中,放量突破在即")。结尾用【操作建议:观望 / 轻仓试探 / 逢低吸纳 / 持有 / 减仓 / 规避】给出明确倾向。
### 2. 📈 技术面分析(核心维度)
这是你的主战场,务必深入:
- **趋势判断**:均线多头/空头排列、20/60 日均线方向、价格在均线之上/下
- **形态结构**:近期是否有突破/破位/双底/双顶/旗形等关键形态
- **指标信号**:MACD 金叉/死叉/背离、KDJ 超买超卖、RSI 强弱、布林通道位置
- **量价配合**:放量上涨/缩量回调/量价背离/换手率异动
每条结论必须引用具体数值(如"MACD 在 6/12 出现金叉,DIF 0.32 上穿 DEA 0.18")。
### 3. 💰 关键价位(买卖区间)
基于提供的关键价位数据,明确指出:
- **上方压力位**(逐档列出,标注强度):第一压力、第二压力
- **下方支撑位**(逐档列出,标注强度):第一支撑、第二支撑
- 给出**建议买入区间**与**止损位**(基于支撑位)
用数据说话,引用提供的压力/支撑(成交密集区)/枢轴点数值。
### 4. 🏭 基本面与财务面(辅助验证)
简要点评(2-4 句,不展开长篇):
- 盈利质量(ROE / 毛利率水平)、成长性(营收/利润增速)
- 与技术面的**交叉验证**:好公司 + 技术面走坏 → 仍需谨慎;差公司 + 技术面强势 → 警惕炒作风险
**当用户消息中标注了"该标的暂无财务数据"时**,本节请输出:
> 📌 财务面分析能力正在接入中。当前版本(Free)未同步该标的的财务报表,基本面维度暂无法评估。
> 技术面分析不依赖财务数据,以下结论依然有效;升级套餐或等待财务数据同步后可补充本维度。
**绝对不要**在无数据时编造 ROE / 增速等数字。
### 5. 📰 消息面(价量异动推断)
**注意:本期无直接新闻数据输入。** 请基于 K 线的**异动信号**进行推断(如:
- 涨停/连板/炸板 → 可能有利好或资金炒作
- 放量暴跌 → 可能有未公开利空
- 突破放量 → 可能有催化剂
明确标注"[推断]",告诉用户这是基于价量的推测,真实消息面数据待接入。若无明显异动,直说"近期价量平稳,无明显消息面信号"
### 6. ⚖️ 综合研判与操作建议
2-3 段:
- 该股当前处于(底部启动 / 上升途中 / 高位震荡 / 下跌趋势 / 底部企稳)哪个阶段
- 风险收益比评估(距支撑位的空间 vs 距压力位的空间)
- **明确操作建议**:激进型 / 稳健型 / 保守型 分别怎么应对
- **需要重点盯的信号**(如跌破 X 支撑止损、站上 Y 压力加仓)
## 分析准则(务必遵守)
1. **技术面优先**:作为交易员,技术面和量价是主要依据,基本面是验证手段,主次分明
2. **数据说话**:每个判断引用具体数值,严禁空泛套话("走势良好"必须改成"连续 3 日站稳 20 日均线且放量")
3. **诚实中立**:看多就写多,看空就写空,不要模棱两可骑墙;数据不支持时直言无法判断
4. **价位精确**:买卖区间必须落到具体价格,基于提供的关键价位数据推演
5. **风险前置**:任何买入建议都要配止损位;提示潜在风险不回避
6. **简明实战**:用交易员能扫读的密度输出,总字数 1000-1800 字,重在可执行
## 重要免责
报告末尾附一行:"> ⚠️ 本报告由 AI 基于公开行情与财务数据生成,仅供参考,不构成任何投资建议。交易有风险,入市需谨慎。"
现在请基于下方数据进行分析。"""
# ================================================================
# 用户消息构建
# ================================================================
def _build_user_prompt(
kline_tail: list[dict],
fins: dict[str, list[dict]],
levels: dict[str, list[dict]],
close: float | None,
symbol: str,
focus: str,
) -> str:
"""构建用户消息:标的 + 价位摘要 + 技术指标 JSON + 财务摘要 + 关注点。"""
parts: list[str] = [
f"标的标准代码: {symbol}",
f"关键价位概览: {summarize_levels(levels, close)}",
"",
"以下是该标的最近日 K 数据(JSON,含 OHLCV 与已计算的技术指标。"
f"最近 {_KLINE_WINDOW} 个交易日,升序):",
"```json",
json.dumps(kline_tail, ensure_ascii=False),
"```",
]
has_fin = any(fins.values())
if has_fin:
parts.extend([
"",
"以下是该标的最新财务数据(JSON,核心指标 + 利润表,金额单位为元):",
"```json",
json.dumps(fins, ensure_ascii=False),
"```",
])
else:
parts.extend([
"",
"(该标的暂无财务数据:当前为 Free 模式或尚未同步财务报表。"
"请按系统提示词第 4 节的说明,在基本面/财务面维度给出\"接入中\"的友好提示,不要编造数据。)",
])
if focus.strip():
parts.extend(["", f"本次分析请特别关注: {focus.strip()}"])
return "\n".join(parts)
# ================================================================
# 关键列筛选(控制上下文体积)
# ================================================================
_KLINE_KEEP_COLS = [
"date", "open", "high", "low", "close", "volume", "change_pct",
"ma5", "ma10", "ma20", "ma60",
"macd_dif", "macd_dea", "macd_hist",
"kdj_k", "kdj_d", "kdj_j",
"rsi_6", "rsi_14", "rsi_24",
"boll_upper", "boll_mid", "boll_lower",
"atr_14", "vol_ratio_5d", "turnover_rate",
"consecutive_limit_ups",
# 信号类(布尔)——只挑对消息面推断有用的几个
"signal_limit_up", "signal_broken_limit_up", "signal_macd_golden",
"signal_macd_death", "signal_ma_golden_5_20", "signal_volume_surge",
"signal_boll_breakout_upper", "signal_boll_breakout_lower",
]
# ================================================================
# 流式分析入口
# ================================================================
async def analyze_stock_stream(
repo,
data_dir: Path,
symbol: str,
focus: str = "",
) -> AsyncIterator[str]:
"""流式个股分析:yield 出每个 NDJSON 事件。
协议(与 financial_analyzer 一致,前端解析无差异):
{"type":"meta","symbol","summary","levels"} 数据 + 价位摘要
{"type":"delta","content":"..."} 逐 chunk 文本
{"type":"error","message":"..."}
{"type":"done"}
"""
# 1. 加载 K 线
df = _load_kline(repo, symbol)
if df.is_empty():
yield json.dumps({
"type": "error",
"message": f"标的 {symbol} 暂无日 K 数据,请先同步",
}, ensure_ascii=False)
return
# 2. 价位计算(基于 K 线)
levels = compute_levels(df)
close = float(df.tail(1)["close"][0]) if "close" in df.columns else None
# 3. 财务(辅助)
fins = _load_financials(data_dir, symbol)
# 4. meta
yield json.dumps({
"type": "meta",
"symbol": symbol,
"summary": summarize_levels(levels, close),
"levels": levels,
"close": close,
}, ensure_ascii=False)
# 5+6. 构建提示词 + 流式调用 LLM(整体 try-except,任何异常都 yield error,避免前端卡死)
try:
from app.services.ai_provider import stream_ai_text
kline_tail = _clean_rows(df, _KLINE_KEEP_COLS)
user_prompt = _build_user_prompt(kline_tail, fins, levels, close, symbol, focus)
async for delta in stream_ai_text(
[
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": user_prompt},
],
temperature=0.5,
max_tokens=4500,
):
yield json.dumps({"type": "delta", "content": delta}, ensure_ascii=False)
except Exception as e: # noqa: BLE001
logger.exception("AI stock analysis failed for %s: %s", symbol, e)
yield json.dumps({"type": "error", "message": f"AI 分析失败: {e}"}, ensure_ascii=False)
return
yield json.dumps({"type": "done"}, ensure_ascii=False)
@@ -0,0 +1,91 @@
"""AI 个股分析报告持久化存储。
与 ai_reports.py(财务分析报告)完全独立 —— 单独的文件、字段、上限,
互不影响。刻意不复用,避免引入 kind 判别字段与分支(解耦 > 抽象)。
存储位置: data/user_data/ai_stock_reports.json (数组,按 created_at 降序)
保留最近 MAX_REPORTS 条;超出自动裁剪最旧的。
每条报告结构:
{
"id": "sar_xxx", # 唯一 id(stock-analysis-report)
"symbol": "600519.SH",
"name": "贵州茅台",
"focus": "", # 用户追加的关心点(可为空)
"content": "# ...markdown", # 报告正文
"summary": "当前价 12.3 · 压力位...", # 价位/数据摘要
"levels": {...}, # 报告生成时的关键价位(供图表回放)
"close": 12.3, # 报告生成时的收盘价
"created_at": "2026-06-26T10:00:00"
}
"""
from __future__ import annotations
import json
import logging
import time
from pathlib import Path
logger = logging.getLogger(__name__)
MAX_REPORTS = 50
def _path() -> Path:
from app.config import settings
p = settings.data_dir / "user_data" / "ai_stock_reports.json"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def list_reports() -> list[dict]:
"""返回全部报告(按 created_at 降序)。"""
p = _path()
if not p.exists():
return []
try:
data = json.loads(p.read_text(encoding="utf-8"))
if isinstance(data, list):
return sorted(data, key=lambda r: r.get("created_at", ""), reverse=True)
except Exception as e: # noqa: BLE001
logger.warning("ai_stock_reports.json malformed: %s", e)
return []
def _save_all(reports: list[dict]) -> None:
"""全量写入(裁剪到 MAX_REPORTS)。"""
reports.sort(key=lambda r: r.get("created_at", ""), reverse=True)
if len(reports) > MAX_REPORTS:
reports = reports[:MAX_REPORTS]
_path().write_text(
json.dumps(reports, indent=2, ensure_ascii=False), encoding="utf-8",
)
def save_report(report: dict) -> dict:
"""新增一条报告并持久化。返回保存后的报告(含 id / created_at)。"""
reports = list_reports()
if not report.get("id"):
report["id"] = f"sar_{int(time.time() * 1000)}_{report.get('symbol', 'x')}"
if not report.get("created_at"):
report["created_at"] = _now_iso()
reports.append(report)
_save_all(reports)
logger.info("Stock report saved: %s (%s), total %d", report.get("symbol"), report.get("id"), len(reports))
return report
def delete_report(report_id: str) -> bool:
"""删除指定报告。返回是否删除成功。"""
reports = list_reports()
before = len(reports)
reports = [r for r in reports if r.get("id") != report_id]
if len(reports) < before:
_save_all(reports)
return True
return False
def _now_iso() -> str:
from datetime import datetime
return datetime.now().isoformat(timespec="seconds")
@@ -0,0 +1,149 @@
"""策略结果缓存 — 写入本地文件,供策略页面秒加载。
缓存结构:
{
"as_of": "2024-01-15",
"results": { strategy_id: { total, as_of, rows } },
"today_ever_matched": { strategy_id: [symbol, ...] }, // 今日曾命中 symbol 并集
"today_ever_rows": { strategy_id: { symbol: row_data } },// 今日曾命中的完整行数据
"updated_at": 1705324800000 # Unix ms
}
文件路径: data/user_data/strategy_cache.json
"""
from __future__ import annotations
import json
import logging
import time
from datetime import date, datetime
from pathlib import Path
from typing import Any
def _json_default(obj: Any) -> Any:
"""处理 date/datetime 等 JSON 不认识的类型。"""
if isinstance(obj, date):
return obj.isoformat()
if isinstance(obj, datetime):
return obj.isoformat()
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
logger = logging.getLogger(__name__)
_CACHE_FILENAME = "strategy_cache.json"
def _cache_path(data_dir: Path) -> Path:
return data_dir / "user_data" / _CACHE_FILENAME
def _enriched_parquet_path(data_dir: Path, as_of: str) -> Path:
"""返回 enriched parquet 文件路径。"""
return data_dir / "kline_daily_enriched" / f"date={as_of}" / "part.parquet"
def _get_enriched_mtime(data_dir: Path, as_of: str) -> float | None:
"""返回 enriched parquet 文件的 mtime (秒)。文件不存在返回 None。"""
p = _enriched_parquet_path(data_dir, as_of)
try:
return p.stat().st_mtime
except FileNotFoundError:
return None
def read_cache(data_dir: Path) -> dict | None:
"""读取策略缓存文件。返回 None 表示无缓存或读取失败。
说明: 原先有 enriched mtime 过期校验 (数据文件变化 → 判过期返回 None),
但在有实时行情的系统里, enriched parquet 每轮被刷新 → mtime 必然变化 →
缓存被永久判死, 策略页读不到数据。且判过期后不触发重算, 只能让用户手动重跑,
保护价值有限。故移除: 盘后缓存总能读出, 实时新鲜度由 /api/screener/cached
端点叠加监控引擎的内存实时结果 (latest_strategy_results) 来保证。
"""
path = _cache_path(data_dir)
if not path.exists():
return None
try:
text = path.read_text(encoding="utf-8")
if not text.strip():
return None
cached = json.loads(text)
except Exception as e: # noqa: BLE001
logger.warning("读取策略缓存失败: %s", e)
return None
return cached
def _rows_to_symbol_map(rows: list[dict]) -> dict[str, dict]:
"""将 rows 列表转为 {symbol: row_data} 映射。"""
result: dict[str, dict] = {}
for row in rows:
sym = row.get("symbol")
if sym:
result[sym] = row
return result
def write_cache(
data_dir: Path,
as_of: str,
results: dict[str, Any],
) -> None:
"""将策略结果写入缓存文件,同时更新今日曾命中集合。
- 日期变更时重置 today_ever_matched 和 today_ever_rows
- 同一天内合并 (并集) 之前曾命中的 symbol,并用最新行数据更新
"""
path = _cache_path(data_dir)
path.parent.mkdir(parents=True, exist_ok=True)
# 读取旧缓存
old = read_cache(data_dir)
old_as_of = old.get("as_of") if old else None
old_ever_rows: dict[str, dict[str, dict]] = old.get("today_ever_rows", {}) if old else {}
# 当前命中的行数据 → symbol 映射
current_row_maps: dict[str, dict[str, dict]] = {}
for sid, r in results.items():
current_row_maps[sid] = _rows_to_symbol_map(r.get("rows", []))
if old_as_of and old_as_of == as_of and old_ever_rows:
# 同一天: 合并 — 用当前行数据更新旧数据 (保持最新价格等)
merged_rows: dict[str, dict[str, dict]] = {}
all_keys = set(old_ever_rows.keys()) | set(current_row_maps.keys())
for sid in all_keys:
old_map = old_ever_rows.get(sid, {})
cur_map = current_row_maps.get(sid, {})
# 以旧数据为基础,用当前数据覆盖 (当前数据更新鲜)
combined = {**old_map, **cur_map}
merged_rows[sid] = combined
today_ever_rows = merged_rows
else:
# 新的一天或首次写入
today_ever_rows = current_row_maps
# 从 ever_rows 提取 symbol 列表 (用于快速计数)
today_ever_matched = {sid: sorted(maps.keys()) for sid, maps in today_ever_rows.items()}
# enriched_mtime: 盘后缓存写入时记录 (向后兼容旧字段)。read_cache 已不再用它
# 做过期校验, 实时新鲜度改由 /cached 端点叠加监控引擎内存结果保证。
enriched_mtime = _get_enriched_mtime(data_dir, as_of)
payload = {
"as_of": as_of,
"results": results,
"today_ever_matched": today_ever_matched,
"today_ever_rows": today_ever_rows,
"enriched_mtime": enriched_mtime,
"updated_at": int(time.time() * 1000),
}
try:
path.write_text(json.dumps(payload, ensure_ascii=False, default=_json_default), encoding="utf-8")
total_rows = sum(len(r.get("rows", [])) for r in results.values())
total_ever = sum(len(v) for v in today_ever_matched.values())
logger.info("策略缓存已写入: %s, %d 策略, %d 命中, %d 曾命中", as_of, len(results), total_rows, total_ever)
except Exception as e: # noqa: BLE001
logger.warning("写入策略缓存失败: %s", e)
+144
View File
@@ -0,0 +1,144 @@
"""自选股服务(§6.1)。
存储:`data/user_data/watchlist.parquet`,字段 symbol + added_at + note。
"""
from __future__ import annotations
import logging
from datetime import datetime
from pathlib import Path
import polars as pl
from app.config import settings
from app.tickflow.capabilities import Cap, CapabilitySet
from app.tickflow.client import get_client
logger = logging.getLogger(__name__)
def _path() -> Path:
p = settings.data_dir / "user_data" / "watchlist.parquet"
p.parent.mkdir(parents=True, exist_ok=True)
return p
def list_symbols() -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
if df.is_empty():
return []
return df.to_dicts()
def add(symbol: str, note: str = "") -> list[dict]:
p = _path()
if p.exists():
df = pl.read_parquet(p)
# 已存在则先移除,后面重新插入到最前面
if symbol in df["symbol"].to_list():
df = df.filter(pl.col("symbol") != symbol)
else:
df = pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8})
new_row = pl.DataFrame({
"symbol": [symbol],
"added_at": [datetime.utcnow().isoformat(timespec="seconds")],
"note": [note],
})
out = pl.concat([new_row, df], how="diagonal_relaxed")
out.write_parquet(p)
return out.to_dicts()
def remove(symbol: str) -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
df = df.filter(pl.col("symbol") != symbol)
df.write_parquet(p)
return df.to_dicts()
def move_to_top(symbol: str) -> list[dict]:
p = _path()
if not p.exists():
return []
df = pl.read_parquet(p)
if df.is_empty() or symbol not in df["symbol"].to_list():
return df.to_dicts()
target = df.filter(pl.col("symbol") == symbol)
rest = df.filter(pl.col("symbol") != symbol)
out = pl.concat([target, rest], how="diagonal_relaxed")
out.write_parquet(p)
return out.to_dicts()
def clear() -> int:
"""清空自选列表。返回移除的数量。"""
p = _path()
if not p.exists():
return 0
df = pl.read_parquet(p)
count = df.height
if count > 0:
pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8}).write_parquet(p)
return count
def fetch_quotes(symbols: list[str], capset: CapabilitySet, timeout_s: float = 8.0) -> list[dict]:
"""拉取实时行情。
优先用 quote.batch;否则降级为 quote.by_symbol 单股请求。
timeout_s: 单批次请求超时(秒),防止 API 卡死阻塞整个请求。
"""
from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeout
if not symbols:
return []
tf = get_client()
quotes: list[dict] = []
# 走 batch
batch_size = 5
if capset.has(Cap.QUOTE_BATCH):
lim = capset.limits(Cap.QUOTE_BATCH)
batch_size = lim.batch if lim and lim.batch else 50
elif capset.has(Cap.QUOTE_BY_SYMBOL):
lim = capset.limits(Cap.QUOTE_BY_SYMBOL)
batch_size = lim.batch if lim and lim.batch else 5
else:
# 无任何实时行情能力(none/free 档走 free-api 服务器,不提供实时行情)
# 提前返回空,避免发起注定失败的请求
return []
chunks = [symbols[i:i + batch_size] for i in range(0, len(symbols), batch_size)]
# 用线程池为每个批次加超时保护
pool = ThreadPoolExecutor(max_workers=1)
for chunk in chunks:
try:
future = pool.submit(tf.quotes.get, symbols=chunk, as_dataframe=True)
raw = future.result(timeout=timeout_s)
if raw is None or len(raw) == 0:
continue
df = pl.from_pandas(raw)
rename_map = {
"last_price": "price",
"ext.change_pct": "pct",
"ext.name": "name",
}
df = df.rename({k: v for k, v in rename_map.items() if k in df.columns})
quotes.extend(df.to_dicts())
except FuturesTimeout:
logger.warning("quote fetch timeout (%.1fs) for %d symbols", timeout_s, len(chunk))
break # 超时后不再尝试后续批次
except Exception as e: # noqa: BLE001
logger.warning("quote fetch failed for %d symbols: %s", len(chunk), e)
pool.shutdown(wait=False)
return quotes
@@ -0,0 +1,170 @@
"""Webhook 推送适配器 — 把告警事件推送到外部 IM / 量化软件。
职责: 把后端产生的告警事件, 通过用户配置的 Webhook 地址推送到外部。
目前支持飞书群机器人; QMT / ptrade 等量化通道为待定。
飞书自定义机器人接入:
1. 飞书群 → 群设置 → 群机器人 → 添加「自定义机器人」
2. 复制生成的 Webhook 地址 (形如 https://open.feishu.cn/open-apis/bot/v2/hook/xxx)
3. (可选) 安全设置 → 启用「签名校验」, 记录签名密钥(secret)
4. 填入设置页「飞书 Webhook」配置
设计: 失败静默降级, 绝不因推送失败阻断告警主流程 (落盘 / SSE 推送)。
去重不在本层做, 复用 MonitorRuleEngine 的 cooldown。
"""
from __future__ import annotations
import base64
import hashlib
import hmac
import logging
import time
logger = logging.getLogger(__name__)
# 单次推送最长字符 (飞书单条文本消息上限 30KB, 这里保守截断避免刷屏)
_MAX_LEN = 500
# 卡片消息正文最长字符 (飞书 interactive 卡片上限 30KB, 保守留余量给标题/结构)
_CARD_MAX_LEN = 28000
# 飞书自定义机器人 Webhook 前缀 (用于 URL 合法性校验)
FEISHU_HOOK_PREFIX = "https://open.feishu.cn/open-apis/bot/v2/hook/"
def _truncate(text: str) -> str:
"""截断超长文本。"""
text = (text or "").strip()
return text[:_MAX_LEN] + ("" if len(text) > _MAX_LEN else "")
def is_valid_feishu_url(url: str) -> bool:
"""校验是否为合法的飞书自定义机器人 Webhook 地址。"""
return bool(url) and url.startswith(FEISHU_HOOK_PREFIX)
def _gen_sign(timestamp: str, secret: str) -> str:
"""计算飞书自定义机器人签名。
算法 (官方): 把 `timestamp + "\\n" + secret` 作为签名字符串 (key),
用 HmacSHA256 计算空字符串的签名结果, 再 Base64 编码。
"""
string_to_sign = f"{timestamp}\n{secret}"
hmac_code = hmac.new(
string_to_sign.encode("utf-8"),
digestmod=hashlib.sha256,
).digest()
return base64.b64encode(hmac_code).decode("utf-8")
def _truncate_card(text: str) -> str:
"""截断卡片正文 (留余量给标题与卡片结构)。"""
text = (text or "").strip()
return text[:_CARD_MAX_LEN] + ("" if len(text) > _CARD_MAX_LEN else "")
def _post_feishu(webhook_url: str, payload: dict, secret: str) -> bool:
"""发送一次飞书 webhook 请求并判定成败 (供 text / card 共用)。
成功响应: HTTP 200 且业务 code=0 (或非 JSON 的 200)。失败静默返回 False。
"""
try:
import httpx
# 启用签名校验时, 请求体须带 timestamp + sign (秒级时间戳)
if secret:
timestamp = str(int(time.time()))
payload["timestamp"] = timestamp
payload["sign"] = _gen_sign(timestamp, secret)
resp = httpx.post(webhook_url, json=payload, timeout=5.0)
# 飞书成功响应: {"code":0,"msg":"success"} (或 StatusCode 200 + Extra)
if resp.status_code == 200:
try:
data = resp.json()
# code=0 表示飞书业务侧成功; 部分版本无 code 字段则按 msg 判断
if isinstance(data, dict):
code = data.get("code", data.get("StatusCode", 0))
if code == 0:
return True
logger.debug("飞书推送业务失败: %s", data)
return False
except ValueError:
# 非 JSON 响应但 HTTP 200, 视为成功
return True
logger.debug("飞书推送 HTTP %s: %s", resp.status_code, resp.text[:200])
return False
except Exception as e: # noqa: BLE001
logger.debug("飞书 Webhook 推送失败: %s", e)
return False
def send_feishu(webhook_url: str, title: str, body: str, secret: str = "") -> bool:
"""推送一条文本消息到飞书群机器人。
Args:
webhook_url: 飞书自定义机器人 Webhook 地址
title: 消息标题 (与正文拼接为一条文本)
body: 消息正文
secret: 签名密钥 (机器人启用了「签名校验」时必填; 留空则不带签名)
Returns:
True=成功送达, False=失败或 URL 非法。
失败静默, 不抛异常 (Webhook 是辅助通道, 不能阻断告警主流程)。
"""
if not is_valid_feishu_url(webhook_url):
return False
text = _truncate(f"{title}\n{body}".strip())
if not text:
return False
payload: dict = {"msg_type": "text", "content": {"text": text}}
return _post_feishu(webhook_url, payload, secret)
def send_feishu_card(webhook_url: str, title: str, subtitle: str, body_md: str, secret: str = "") -> bool:
"""推送一条 interactive 卡片消息到飞书群机器人 —— 用 lark_md 渲染完整 markdown 报告。
飞书「自定义机器人」webhook 不支持文件附件, 但 interactive 卡片的 lark_md 元素
可渲染 markdown, 能承载完整复盘报告(通常 2-5KB, 远小于卡片 30KB 上限)。
Args:
webhook_url: 飞书自定义机器人 Webhook 地址
title: 卡片标题 (显示在蓝色 header)
subtitle: 副标题 (加粗显示, 如日期/情绪标签; 留空则省略)
body_md: 卡片正文 markdown (报告全文)
secret: 签名密钥 (启用签名校验时必填)
Returns:
True=成功送达, False=失败或 URL 非法。
失败静默, 不抛异常 (与 send_feishu 一致, 不阻断告警主流程)。
"""
if not is_valid_feishu_url(webhook_url):
return False
body = _truncate_card(body_md)
elements: list[dict] = []
if subtitle.strip():
elements.append({
"tag": "div",
"text": {"tag": "lark_md", "content": f"**{subtitle.strip()}**"},
})
elements.append({"tag": "hr"})
elements.append({
"tag": "div",
"text": {"tag": "lark_md", "content": body},
})
payload: dict = {
"msg_type": "interactive",
"card": {
"config": {"wide_screen_mode": True},
"header": {
"title": {"tag": "plain_text", "content": title},
"template": "blue",
},
"elements": elements,
},
}
return _post_feishu(webhook_url, payload, secret)