release: v1.29.0 — 借鉴社区 Fork 六项特性:ZIG 策略 + 交易时段感知刷新 + 120M K 线 + 逐 bar 衍生字段 + 159 龙头池 + 多 Provider LLM 直连

- ZIG 右侧突破回补策略:MyTT 新增 ZIG 之字转向(未来函数,含前视偏差警示);
  波谷启动建仓挂硬止损(OCO)→ 见顶清仓记前高 → 右侧突破回补;路径依赖不实现
  entry_exit_masks(向量化守护测试白名单);含寻优预设网格与 --strategy-file 独立文件
- 交易时段感知刷新:realtime/session.py(09:15~11:30:30 / 13:00~15:05)+
  GET /market/session;看板 30/60/120s 轮询休市自动暂停(三态状态栏 + 开关持久化
  + 手动刷新不受限);SSE/WS 既有会话语义不动
- 120 分钟 K 线:/bars?category=MIN_120(MAC 原生 Period.MINS×120 优先,
  2×60M 相邻聚合兜底,标准客户端上限 400 根);前端周期选择器同步
- 逐 bar 衍生字段:/bars 与 /bars/index 附带 pre_close/change/change_pct/
  amplitude_pct(pre_close≤0.01 兜底防除零)
- 159 只核心龙头池:数据资产取自 Fork(东财全行业龙头名单,四组分层);
  universe=core 接入 screen scan / SignalScanner / StrengthRanker / market strength;
  GET /market/core-leaders + WebUI「龙头池」页(搜索/个股详情)
- 多 Provider LLM 直连:easy_tdx.ai + /llm/*(DeepSeek/通义/智谱/Kimi/MiniMax/
  OpenAI/Claude/Ollama/自定义,openai 兼容 + anthropic 原生双协议);
  配置落盘 ~/.easy_tdx/llm.json(WebUI「AI 设置」页 ⇆ 手工编辑双向兼容,
  文件>环境变量>预设;key 脱敏回显/CLEAR 清除);「AI 解读」后台任务化
  (复用 task_runner,提交+轮询,不占 HTTP 连接);思考型模型空白正文防御
  (reasoning_content 耗尽 max_tokens → 可操作报错;默认 16000);
  AI 解读历史页(自动归档 Prompt/正文/策略上下文 + 去回测带参引导)
- WebUI 加固:SPA fallback 对未知 /api/* 返回 JSON 404(不再 200 HTML 伪装解析错);
  index.html 一律 Cache-Control: no-store(防缓存旧资源引用);路由兜底重定向;
  全局风险提示常驻底栏 + 龙头池/AI 解读针对性免责声明
- 测试:新增 9 个单测文件共 59 例;黄金基线仅新增 zig_breakout 条目(其余零漂移);
  全量 1448 例通过
This commit is contained in:
GitHub
2026-09-02 20:19:17 +08:00
parent eeed45b171
commit 4bd5b5d833
45 changed files with 4275 additions and 70 deletions
+113
View File
@@ -571,4 +571,117 @@ def FSL(CLOSE, VOL, CAPITAL): # 分水岭指标:多空趋势强弱分界(SW
return RD(SWL), RD(SWS)
def ZIG(S, X=35): # 之字转向指标(未来函数):S为价格序列,X为转向阈值百分比(如10表示10%)
"""之字转向指标 (ZigZag) — 经典未来函数。
当价格从前一个极值点反向变动超过 X% 时确立波峰/波谷拐点并转向,
拐点之间线性插值,返回与 S 等长的拟合序列。
注意:拐点只有在**其后**的走势确认了转向才会回溯标出,序列中波峰/
波谷位置含有未来信息。把 ZIG 拐点直接当买卖信号回测会严重高估收益
(前视偏差);如需使用,必须配合右侧确认或止损保护(参见内置策略
``zig_breakout`` 的做法)。
Args:
S: 价格序列(通常为 CLOSE)
X: 转向阈值百分比。10 表示 10%;也可传小数形式 0.1(以 1.0 为界
自动区分,故阈值本身小于 1% 时请用小数形式)
Returns:
np.ndarray: 与 S 等长的 ZIG 之字转向插值序列
"""
S = np.asarray(S, dtype=float)
n = len(S)
if n == 0:
return np.array([], dtype=float)
if n == 1:
return S.copy()
x = float(X) / 100.0 if float(X) > 1.0 else float(X)
if x <= 0:
return S.copy()
ZIG_STATE_START = 0
ZIG_STATE_RISE = 1
ZIG_STATE_FALL = 2
peer_i = 0
candidate_i = None
peers = [0]
state = ZIG_STATE_START
for scan_i in range(1, n):
if scan_i == n - 1:
# 扫描到序列尾部:未确立的候选极值按当前方向收尾
if candidate_i is None:
peers.append(scan_i)
else:
if state == ZIG_STATE_RISE:
if S[scan_i] >= S[candidate_i]:
peers.append(scan_i)
else:
peers.append(candidate_i)
if candidate_i != scan_i:
peers.append(scan_i)
elif state == ZIG_STATE_FALL:
if S[scan_i] <= S[candidate_i]:
peers.append(scan_i)
else:
peers.append(candidate_i)
if candidate_i != scan_i:
peers.append(scan_i)
else:
peers.append(scan_i)
break
if state == ZIG_STATE_START:
if S[peer_i] != 0:
if S[scan_i] >= S[peer_i] * (1.0 + x):
candidate_i = scan_i
state = ZIG_STATE_RISE
elif S[scan_i] <= S[peer_i] * (1.0 - x):
candidate_i = scan_i
state = ZIG_STATE_FALL
elif state == ZIG_STATE_RISE:
if S[scan_i] >= S[candidate_i]:
candidate_i = scan_i
elif S[candidate_i] != 0 and S[scan_i] <= S[candidate_i] * (1.0 - x):
peer_i = candidate_i
peers.append(peer_i)
state = ZIG_STATE_FALL
candidate_i = scan_i
elif state == ZIG_STATE_FALL:
if S[scan_i] <= S[candidate_i]:
candidate_i = scan_i
elif S[candidate_i] != 0 and S[scan_i] >= S[candidate_i] * (1.0 + x):
peer_i = candidate_i
peers.append(peer_i)
state = ZIG_STATE_RISE
candidate_i = scan_i
# 去除重复拐点并确保末端对齐
clean_peers = []
for p in peers:
if not clean_peers or p != clean_peers[-1]:
clean_peers.append(p)
if clean_peers[-1] != n - 1:
clean_peers.append(n - 1)
# 拐点间线性插值
z = np.zeros(n, dtype=float)
for i in range(len(clean_peers) - 1):
p_start = clean_peers[i]
p_end = clean_peers[i + 1]
v_start = S[p_start]
v_end = S[p_end]
if p_end == p_start:
z[p_start] = v_start
else:
slope = (v_end - v_start) / (p_end - p_start)
for j in range(p_end - p_start + 1):
z[p_start + j] = v_start + slope * j
return RD(z)
# 望大家能提交更多指标和函数 https://github.com/mpquant/MyTT
+1
View File
@@ -110,6 +110,7 @@ def FSL(
VOL: npt.ArrayLike,
CAPITAL: float,
) -> tuple[NDArray, NDArray]: ...
def ZIG(S: npt.ArrayLike, X: float = ...) -> NDArray: ...
# ── Utility Functions ────────────────────────────────────────────────────────
+35
View File
@@ -0,0 +1,35 @@
"""LLM 客户端与配置(多 Provider,WebUI 与配置文件双向兼容)。
借鉴社区 Forkswimmingaaron/easy_tdx)的极简 LLM 客户端思路并扩展:
国产主流 Provider 预设(DeepSeek/通义千问/智谱/Kimi/MiniMax+ OpenAI/
Claude/Ollama + 完全自定义,统一收敛到两种线上协议(openai 兼容 /
anthropic 原生)。
配置来源(字段级优先级,WebUI 与配置文件天然兼容)::
~/.easy_tdx/llm.json 字段值 > 环境变量 > Provider 预设默认值
WebUI 保存 = 写这个 JSON 文件;手工编辑文件 = 下次请求即生效。环境变量
``LLM_PROVIDER`` / ``LLM_API_KEY`` / ``LLM_BASE_URL`` / ``LLM_MODEL``
与常见工具惯例一致,仅在文件缺字段时兜底。
"""
from easy_tdx.ai.llm import (
PROVIDER_PRESETS,
LlmClient,
LlmConfig,
load_config,
mask_key,
resolve_config,
save_config,
)
__all__ = [
"PROVIDER_PRESETS",
"LlmClient",
"LlmConfig",
"load_config",
"mask_key",
"resolve_config",
"save_config",
]
+401
View File
@@ -0,0 +1,401 @@
"""多 Provider LLM 客户端:配置解析 + HTTP 调用。
零第三方依赖:HTTP 走标准库 urllib(经 ``asyncio.to_thread`` 异步化),
FastAPI 路由可直接 ``await``。
Provider 预设(``api_style``):
=========== ======== ============================================== ==================
provider 协议 base_url 默认模型
=========== ======== ============================================== ==================
deepseek openai https://api.deepseek.com/v1 deepseek-chat
qwen openai https://dashscope.aliyuncs.com/compatible-mode qwen-plus
/v1
zhipu openai https://open.bigmodel.cn/api/paas/v4 glm-4-flash
kimi openai https://api.moonshot.cn/v1 moonshot-v1-8k
minimax openai https://api.minimaxi.chat/v1 MiniMax-Text-01
openai openai https://api.openai.com/v1 gpt-4o-mini
claude anthropic https://api.anthropic.com/v1 claude-sonnet-4-5
ollama openai http://localhost:11434/v1 qwen2.5:7b
custom openai (用户填写) (用户填写)
=========== ======== ============================================== ==================
预设的 base_url/默认模型只是初始填充值——WebUI 或 JSON 文件里均可覆盖
(自定义网关/代理场景直接改 url 即可)。
"""
from __future__ import annotations
import asyncio
import json
import logging
import os
import time
import urllib.error
import urllib.request
from dataclasses import asdict, dataclass, field, replace
from pathlib import Path
from typing import Any
__all__ = [
"PROVIDER_PRESETS",
"LlmClient",
"LlmConfig",
"load_config",
"mask_key",
"resolve_config",
"save_config",
]
logger = logging.getLogger(__name__)
#: 配置文件名(落在 EASY_TDX_CONFIG_DIR,与 watchlist/strategies 同目录)。
LLM_CONFIG_FILENAME = "llm.json"
@dataclass
class ProviderPreset:
"""单个 Provider 的展示信息与默认填充值。"""
id: str
label: str
base_url: str
default_model: str
api_style: str = "openai" # "openai" | "anthropic"
needs_key: bool = True # ollama 本地服务无需 key
def to_dict(self) -> dict[str, Any]:
return {
"id": self.id,
"label": self.label,
"base_url": self.base_url,
"default_model": self.default_model,
"api_style": self.api_style,
"needs_key": self.needs_key,
}
#: Provider 预设表(WebUI 下拉框数据源 + 未配置字段的兜底默认值)。
PROVIDER_PRESETS: dict[str, ProviderPreset] = {
p.id: p
for p in (
ProviderPreset("deepseek", "DeepSeek", "https://api.deepseek.com/v1", "deepseek-chat"),
ProviderPreset(
"qwen",
"通义千问 Qwen",
"https://dashscope.aliyuncs.com/compatible-mode/v1",
"qwen-plus",
),
ProviderPreset("zhipu", "智谱 GLM", "https://open.bigmodel.cn/api/paas/v4", "glm-4-flash"),
ProviderPreset("kimi", "Kimi (月之暗面)", "https://api.moonshot.cn/v1", "moonshot-v1-8k"),
ProviderPreset("minimax", "MiniMax", "https://api.minimaxi.chat/v1", "MiniMax-Text-01"),
ProviderPreset("openai", "OpenAI", "https://api.openai.com/v1", "gpt-4o-mini"),
ProviderPreset(
"claude",
"Claude (Anthropic)",
"https://api.anthropic.com/v1",
"claude-sonnet-4-5",
api_style="anthropic",
),
ProviderPreset(
"ollama",
"Ollama(本地)",
"http://localhost:11434/v1",
"qwen2.5:7b",
needs_key=False,
),
ProviderPreset("custom", "自定义(OpenAI 兼容)", "", ""),
)
}
@dataclass
class LlmConfig:
"""LLM 调用配置(WebUI 表单与 llm.json 的公共结构)。"""
provider: str = "deepseek"
api_url: str = "" # 留空 = 用预设 base_url
api_key: str = ""
model: str = "" # 留空 = 用预设默认模型
temperature: float = 0.3
# max_tokens 是"上限"而非目标(按实际生成计费):思考型模型的思考链
# 计入该预算,4000 会被整份报告的思考轻易耗尽导致正文空白,默认给足
timeout: float = 180.0
max_tokens: int = 16000
system_prompt: str = field(
default="你是一位严谨的 A 股量化投研分析师,基于给定的数据客观分析,"
"不确定的内容明确说明,不构成投资建议。"
)
def to_dict(self, *, mask_api_key: bool = False) -> dict[str, Any]:
d = asdict(self)
if mask_api_key:
d["api_key"] = mask_key(self.api_key)
return d
# ── 配置读写(文件 > 环境变量 > 预设) ────────────────────────────────────────
def config_path() -> Path:
"""配置文件路径(``$EASY_TDX_CONFIG_DIR/llm.json``,默认 ``~/.easy_tdx``)。"""
base = Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx")))
return base / LLM_CONFIG_FILENAME
def _read_config_file() -> dict[str, Any]:
"""直读 llm.json(无缓存——手工编辑即时生效)。损坏/不存在返回空 dict。"""
p = config_path()
if not p.is_file():
return {}
try:
data = json.loads(p.read_text(encoding="utf-8"))
return data if isinstance(data, dict) else {}
except (json.JSONDecodeError, OSError) as exc:
logger.warning("读取 LLM 配置失败 %s: %s", p, exc)
return {}
def load_config() -> LlmConfig:
"""加载配置:llm.json 显式字段 > 环境变量兜底(未填字段仍为空,调用时再取预设)。"""
data = _read_config_file()
env_url = os.environ.get("LLM_BASE_URL", "")
cfg = LlmConfig(
provider=str(data.get("provider") or os.environ.get("LLM_PROVIDER", "") or "deepseek"),
api_url=str(data.get("api_url") or env_url or ""),
api_key=str(data.get("api_key") or os.environ.get("LLM_API_KEY", "") or ""),
model=str(data.get("model") or os.environ.get("LLM_MODEL", "") or ""),
temperature=float(data.get("temperature", 0.3)),
max_tokens=int(data.get("max_tokens", 16000)),
timeout=float(data.get("timeout", 180.0)),
system_prompt=str(data.get("system_prompt", "") or LlmConfig.system_prompt),
)
return cfg
def save_config(cfg: LlmConfig) -> Path:
"""写入 llm.json(WebUI 保存入口;目录惰性创建)。"""
p = config_path()
p.parent.mkdir(parents=True, exist_ok=True)
p.write_text(
json.dumps(cfg.to_dict(), ensure_ascii=False, indent=2), encoding="utf-8", newline="\n"
)
return p
def resolve_config(cfg: LlmConfig | None = None) -> LlmConfig:
"""把配置的空字段用 Provider 预设补齐,得到可直接调用的完整配置。
- ``api_url`` 空 → 预设 ``base_url``
- ``model`` 空 → 预设 ``default_model``
- provider 无预设(拼错)→ 按 custom 处理,url/model 必须已填。
Raises:
ValueError: 补齐后仍缺 api_url 或 modelcustom 未填全)。
"""
c = replace(cfg or load_config())
preset = PROVIDER_PRESETS.get(c.provider, PROVIDER_PRESETS["custom"])
if not c.api_url:
c.api_url = preset.base_url
if not c.model:
c.model = preset.default_model
if not c.api_url or not c.model:
raise ValueError(
f"LLM 配置不完整:provider={c.provider} 缺少 api_url 或 model"
"请在 AI 设置中补全"
)
return c
def mask_key(key: str) -> str:
"""API Key 脱敏展示:保头 3 尾 4,中间打码(短 key 全打码)。"""
if not key:
return ""
if len(key) <= 8:
return "*" * len(key)
return f"{key[:3]}***{key[-4:]}"
# ── HTTP 客户端(标准库实现) ─────────────────────────────────────────────────
class LlmError(RuntimeError):
"""LLM 调用失败(网络/鉴权/响应格式)。"""
def __init__(self, message: str, *, status: int | None = None) -> None:
super().__init__(message)
self.status = status
def _post_json(
url: str, headers: dict[str, str], payload: dict[str, Any], timeout: float
) -> dict[str, Any]:
"""同步 POST JSON(在线程池里跑),返回解析后的 JSON。
urllib 默认带 ``User-Agent: Python-urllib``,部分网关拒绝——显式带 UA。
超时单独成类报错:非流式 chat 接口要等模型**整段回复生成完**才回包,
大 Prompt(如整份回测报告解读)生成 1-3 分钟很正常,读超时≠网络故障,
报错必须把「调大超时」这个动作说清楚(v1.29.1 实测踩坑)。
"""
req = urllib.request.Request(
url,
data=json.dumps(payload).encode("utf-8"),
headers={"User-Agent": "easy-tdx/llm", "Content-Type": "application/json", **headers},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
return json.loads(resp.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
body = exc.read().decode("utf-8", errors="replace")[:500]
raise LlmError(f"LLM API HTTP {exc.code}: {body}", status=exc.code) from exc
except urllib.error.URLError as exc:
if isinstance(exc.reason, TimeoutError):
raise LlmError(_timeout_message(timeout)) from exc
raise LlmError(f"LLM API 网络错误: {exc.reason}") from exc
except TimeoutError as exc:
raise LlmError(_timeout_message(timeout)) from exc
except json.JSONDecodeError as exc:
raise LlmError(f"LLM API 响应不是合法 JSON: {exc}") from exc
def _timeout_message(timeout: float) -> str:
return (
f"请求超时({timeout:.0f}s 内无响应)——非流式接口需等模型生成完整段回复,"
"大报告解读 1-3 分钟属正常。可在「AI 设置」调大「超时(秒)」,"
"或换生成更快的模型后重试"
)
class LlmClient:
"""单次配置快照的 LLM 调用客户端(无连接状态,可随时重建)。"""
def __init__(self, cfg: LlmConfig | None = None) -> None:
self._cfg = resolve_config(cfg)
@property
def config(self) -> LlmConfig:
return self._cfg
async def chat(self, prompt: str, system_prompt: str | None = None) -> str:
"""发一轮对话,返回模型回复文本。
Args:
prompt: 用户消息(如回测报告组装成的解读 Prompt)。
system_prompt: 系统提示,None = 用配置里的默认。
Raises:
LlmError: 网络/鉴权/格式错误(含未配置 api_key 的场景)。
"""
cfg = self._cfg
preset = PROVIDER_PRESETS.get(cfg.provider, PROVIDER_PRESETS["custom"])
if preset.needs_key and not cfg.api_key:
raise LlmError(
f"未配置 {preset.label} 的 API Key——请在 WebUI「AI 设置」页"
"或 ~/.easy_tdx/llm.json 中填写(或设置 LLM_API_KEY 环境变量)"
)
system = system_prompt if system_prompt is not None else cfg.system_prompt
return await asyncio.to_thread(self._chat_sync, prompt, system, preset.api_style)
# -- 同步实现(to_thread 里跑) --------------------------------------------
def _chat_sync(self, prompt: str, system: str, api_style: str) -> str:
if api_style == "anthropic":
return self._chat_anthropic(prompt, system)
return self._chat_openai(prompt, system)
def _chat_openai(self, prompt: str, system: str) -> str:
cfg = self._cfg
headers = {"Authorization": f"Bearer {cfg.api_key}"} if cfg.api_key else {}
payload = {
"model": cfg.model,
"messages": [
{"role": "system", "content": system},
{"role": "user", "content": prompt},
],
"temperature": cfg.temperature,
"max_tokens": cfg.max_tokens,
}
url = f"{cfg.api_url.rstrip('/')}/chat/completions"
data = _post_json(url, headers, payload, cfg.timeout)
try:
message = data["choices"][0]["message"]
finish = str(data["choices"][0].get("finish_reason") or "")
return self._extract_reply_openai(message, finish)
except LlmError:
raise
except (KeyError, IndexError, TypeError) as exc:
raw = json.dumps(data, ensure_ascii=False)[:300]
raise LlmError(f"LLM 响应格式异常: {raw}") from exc
def _extract_reply_openai(self, message: dict[str, Any], finish: str) -> str:
"""从 OpenAI 兼容响应的 message 里提取正文,处理思考型模型的空白正文。
思考型模型(GLM-5.x / DeepSeek-R1 / o 系列等)的 ``reasoning_content``
计入 max_tokens:预算被思考链耗尽时 ``content`` 为空白——truthy 但
渲染为空(v1.29.1 实测:状态条报成功、正文空白)。这里显式拦截:
空白正文一律报可操作的错误(提示调大 max_tokens),绝不返回空串。
"""
content = message.get("content")
text = str(content) if content is not None else ""
if text.strip():
return text
reasoning = message.get("reasoning_content") or message.get("reasoning")
if reasoning:
raise LlmError(
f"模型只返回了思考链(reasoning_content {len(str(reasoning))} 字),"
f"未生成正文——max_tokens={self._cfg.max_tokens} 大概率被思考耗尽"
f"finish_reason={finish or 'unknown'})。"
"请在「AI 设置」把 Max Tokens 调大(思考型模型建议 ≥16000)后重试"
)
if finish == "length":
raise LlmError(
"模型输出被 max_tokens 截断且无正文,请在「AI 设置」调大 Max Tokens 后重试"
)
raw = json.dumps(message, ensure_ascii=False)[:300]
raise LlmError(f"LLM 响应 message.content 为空: {raw}")
def _chat_anthropic(self, prompt: str, system: str) -> str:
cfg = self._cfg
headers = {
"x-api-key": cfg.api_key,
"anthropic-version": "2023-06-01",
}
payload = {
"model": cfg.model,
"max_tokens": cfg.max_tokens,
"temperature": cfg.temperature,
"system": system,
"messages": [{"role": "user", "content": prompt}],
}
url = f"{cfg.api_url.rstrip('/')}/messages"
data = _post_json(url, headers, payload, cfg.timeout)
try:
blocks = data["content"]
return "".join(str(b.get("text", "")) for b in blocks if b.get("type") == "text")
except (KeyError, TypeError) as exc:
raw = json.dumps(data, ensure_ascii=False)[:300]
raise LlmError(f"LLM 响应格式异常: {raw}") from exc
async def test(self) -> dict[str, Any]:
"""连通性测试:发一句极短 ping,返回 ok/延迟/样例回复。"""
t0 = time.perf_counter()
try:
reply = await self.chat(
"请只回复两个字:OK", system_prompt="You are a connectivity probe."
)
return {
"ok": True,
"latency_ms": round((time.perf_counter() - t0) * 1000),
"model": self._cfg.model,
"provider": self._cfg.provider,
"reply": reply.strip()[:100],
}
except LlmError as exc:
return {
"ok": False,
"latency_ms": round((time.perf_counter() - t0) * 1000),
"model": self._cfg.model,
"provider": self._cfg.provider,
"error": str(exc),
}
+112
View File
@@ -30,6 +30,7 @@ from easy_tdx.MyTT import (
EMA,
EMV,
FSL,
HHV,
KDJ,
KTN,
MA,
@@ -38,6 +39,7 @@ from easy_tdx.MyTT import (
TAQ,
TRIX,
WR,
ZIG,
)
__all__: list[str] = [] # 注册副作用即可,无需导出符号
@@ -691,3 +693,113 @@ class FslStrategy(ParametrizedStrategy):
def entry_exit_masks(self) -> tuple[Any, Any]:
"""与 next() 同源(gold/dead 即 next() 判定用的同一组掩码数组)。"""
return self.gold, self.dead
# ── ZIG 右侧突破回补 ─────────────────────────────────────────────────────────
@register_strategy(
name="zig_breakout",
label="ZIG 右侧突破回补",
description=(
"ZIG 向上启动(波谷确认)全仓买入;ZIG 见顶回落清仓并记录 N 日最高点,"
"其后收盘突破前高×(1+确认比例) 时右侧回补。两路径买入均带硬止损,"
"对冲 ZIG 波谷确认的前视偏差(未来函数,实盘信号会滞后)。"
),
)
class ZigBreakoutStrategy(ParametrizedStrategy):
"""ZIG 右侧突破回补(Re-entry on Breakout + 硬止损保护)。
ZIG 是未来函数:波峰/波谷只有在其后走势确认转向才回溯标出,回测里
"波谷启动"信号天然偷看未来。本策略用两层保护缓解而非消除该偏差:
1. 买入即挂 ``stop_loss_pct`` 硬止损(引擎逐 bar 监控,跌破自动平仓),
假波谷不至于深套;
2. 卖出后不追 ZIG 新波谷,而是等价格**右侧突破**前高确认_pct 再回补,
"猜底"换成"确认后进场"
交易逻辑::
空仓 + ZIG 上行 → 全仓买入(带止损)
持仓 + ZIG 下行(见顶) → 全仓卖出,记录 HHV(high, N) 为前高
空仓 + 收盘 ≥ 前高×(1+确认) → 右侧回补(带止损)
注意:``_breakout_level`` 随持仓路径变化,信号不可向量化,故不实现
``entry_exit_masks``——引擎自动走逐 bar 回放路径(与 next() 完全一致)。
"""
params = [
Param(
"zig_delta",
float,
default=10.0,
min_value=0.5,
max_value=50.0,
label="ZIG转向阈值%",
),
Param(
"confirm_pct",
float,
default=2.0,
min_value=0.1,
max_value=20.0,
label="突破确认比例%",
),
Param("hhv_period", int, default=20, min_value=5, max_value=120, label="前高周期"),
Param(
"stop_loss_pct",
float,
default=3.0,
min_value=0.0,
max_value=30.0,
label="硬止损%",
),
]
def init(self) -> None:
self.zig = self.I(ZIG, self.data.close, self.p["zig_delta"])
self.hhv = self.I(HHV, self.data.high, self.p["hhv_period"])
# 见顶清仓时记录的前高(0 = 未记录,等待首次建仓-见顶周期)
self._breakout_level: float = 0.0
def next(self) -> None:
i = self._bar_index
if i == 0:
return
cur_close = float(self.data.close[0])
cur_zig = float(self.zig[i])
prev_zig = float(self.zig[i - 1])
cur_pos = self.position["size"]
# 持仓:ZIG 见顶回落 → 清仓,并记录突破位(HHV 含未来 bar 已确认的高点)
if cur_pos > 0 and cur_zig < prev_zig:
self._breakout_level = float(self.hhv[i])
self.sell(size=0)
return
if cur_pos == 0:
# 路径 1:ZIG 向上启动(波谷确认)→ 初始建仓
if cur_zig > prev_zig:
self._breakout_level = 0.0
self._buy_with_stop()
return
# 路径 2:右侧突破前高 → 回补(洗盘结束、主升确立)
if self._breakout_level > 0:
threshold = self._breakout_level * (1.0 + self.p["confirm_pct"] / 100.0)
if cur_close >= threshold:
self._breakout_level = 0.0
self._buy_with_stop()
def _buy_with_stop(self) -> None:
"""市价全仓买入并按 ``stop_loss_pct`` 挂硬止损(0 = 不挂)。
市价单(price=None)由引擎在下一根开盘成交,与本地其他内置策略
口径一致,避免信号 bar 收盘价成交的前视味道。
"""
pct = self.p["stop_loss_pct"] / 100.0
if pct > 0:
self.buy(size=0, stop_loss_pct=pct)
else:
self.buy(size=0)
@@ -98,6 +98,12 @@ STRATEGY_PRESETS: dict[str, dict[str, list[Any]]] = {
# capital 仅作粗档扫描(1千万/1亿/10亿股),覆盖小盘→大盘
"capital": [1e7, 1e8, 1e9, 1e10],
}, # 4
# ── 之字转向类 ───────────────────────────────────────────────────────────
"zig_breakout": {
# 转向阈值×确认比例 = 4×3;hhv/止损用默认(HHV20 / 3%
"zig_delta": [5.0, 8.0, 10.0, 15.0],
"confirm_pct": [1.0, 2.0, 3.0],
}, # 12
}
+4
View File
@@ -23,6 +23,7 @@ from easy_tdx.realtime.engine import (
RealtimeStrategy,
)
from easy_tdx.realtime.feed import RealtimeDataFeed
from easy_tdx.realtime.session import SESSION_WINDOWS, is_trading_time, session_info
__all__ = [
"EventBus",
@@ -31,4 +32,7 @@ __all__ = [
"MarketEvent",
"RealtimeDataFeed",
"RealtimeStrategy",
"SESSION_WINDOWS",
"is_trading_time",
"session_info",
]
+75
View File
@@ -0,0 +1,75 @@
"""A 股交易时段判断(共享工具)。
已有的两处会话过滤各自私有、口径不一:
- :mod:`easy_tdx.realtime.feed` 的 ``_DEFAULT_SESSIONS``09:15-11:30 / 13:00-15:00
WS 按需轮询用,收盘竞价不拉);
- :mod:`easy_tdx.web.quote_streamer` 的 ``_is_trading_hours``09:10-15:10 连续窗,
SSE 快照轮询用,午休也降频拉收盘价快照)。
本模块提供第三个口径——**WebUI 仪表盘自动刷新用的"有效行情时段"**
在 feed 的窗口基础上,早盘前移到 09:15(集合竞价有行情),尾盘后移到
15:05(收盘集合竞价 15:00-15:03 仍有成交),午休排除。前端在此时段内
做 15-30s 轮询,之外暂停自动刷新(手动刷新不受限)。
不改动上述两处既有语义,避免影响它们的测试与行为。
"""
from __future__ import annotations
from datetime import datetime, time, tzinfo
from typing import Any
__all__ = ["SESSION_WINDOWS", "SESSION_DESC", "is_trading_time", "session_info"]
#: 有效行情时段(本地时间)。窗口 = (start, end),含两端。
#: - 早盘 09:15:00-11:30:3009:15 起集合竞价可看,11:30:30 容纳尾单撮合散点;
#: - 午盘 13:00:00-15:05:0015:00-15:03 为收盘集合竞价,留 2 分钟余量。
SESSION_WINDOWS: tuple[tuple[time, time], ...] = (
(time(9, 15, 0), time(11, 30, 30)),
(time(13, 0, 0), time(15, 5, 0)),
)
#: 展示用时段描述(前端状态栏 / API 响应)。
SESSION_DESC = "09:15~11:30, 13:00~15:05"
def is_trading_time(now: datetime | None = None, *, tz: tzinfo | None = None) -> bool:
"""判断当前是否处于 A 股有效行情时段(周一至周五,午休与深夜除外)。
只做"星期 + 时分"判断,不含法定节假日日历——节假日全天处于闭市
窗口外时前端轮询暂停是安全方向(误刷新无副作用,漏刷新才是问题,
而节假日行情本就不动,手动刷新始终可用)。
Args:
now: 待判断时间,None = 取本地当前时间。
tz: 未传 ``now`` 时使用的时区,None = 系统本地时区。
Returns:
True = 盘中(含集合竞价缓冲窗)。
"""
t = now or datetime.now(tz=tz)
if t.weekday() >= 5: # 周六/周日
return False
for start, end in SESSION_WINDOWS:
if start <= t.time() <= end:
return True
return False
def session_info(now: datetime | None = None, *, tz: tzinfo | None = None) -> dict[str, Any]:
"""构建 /market/session 响应体:时段判断 + 窗口描述 + 服务器时间。
前端以本地判断为主(每 15s 重估),本接口用于校准服务器侧视角。
"""
t = now or datetime.now(tz=tz)
return {
"is_trading_time": is_trading_time(t),
"sessions": [
{"start": s.strftime("%H:%M"), "end": e.strftime("%H:%M")}
for s, e in SESSION_WINDOWS
],
"session_desc": SESSION_DESC,
"server_time": t.isoformat(timespec="seconds"),
"weekday": t.weekday(),
}
+3
View File
@@ -24,6 +24,7 @@ from easy_tdx.screen.strength import ( # noqa: F401
StrengthRanker,
StrengthResult,
)
from easy_tdx.screen.universe import CORE_LEADERS, CORE_LEADERS_DESC # noqa: F401
__all__ = [
"SignalScanner",
@@ -31,4 +32,6 @@ __all__ = [
"StrengthRanker",
"StrengthResult",
"STRENGTH_PRESETS",
"CORE_LEADERS",
"CORE_LEADERS_DESC",
]
+1 -1
View File
@@ -35,7 +35,7 @@ def screen() -> None:
@click.option(
"--universe",
default="all",
help="股票范围: all/sh/sz/<文件路径>(默认 all",
help="股票范围: all/sh/sz/core/<文件路径>(默认 allcore=159只核心龙头池",
)
@click.option("--vipdoc", default=None, help="离线数据目录(默认自动检测)")
@click.option("--cash", default=100_000.0, type=float, help="初始资金")
+15 -3
View File
@@ -108,6 +108,7 @@ class SignalScanner:
- "all": 沪深全部 A 股(默认)
- "sh": 仅上海
- "sz": 仅深圳
- "core": 核心龙头池 159 只(跨沪深按名单过滤)
- 文件路径: 每行一个 "市场 代码"(如 "SZ 000001"
progress_callback: 进度回调函数(current, total, filename)
workers: 并发工作进程数
@@ -357,13 +358,20 @@ class SignalScanner:
"""
# 确定要扫描的交易所目录
exchanges: list[str] = []
if universe in ("all", "sz"):
if universe in ("all", "sz", "core"):
exchanges.append("sz")
if universe in ("all", "sh"):
if universe in ("all", "sh", "core"):
exchanges.append("sh")
# 核心龙头池:跨沪深按名单过滤(见 screen/universe.py
core_codes: set[str] | None = None
if universe == "core":
from easy_tdx.screen.universe import core_leader_codes
core_codes = core_leader_codes()
# 从文件列表模式读取
if universe not in ("all", "sh", "sz"):
if universe not in ("all", "sh", "sz", "core"):
return self._collect_from_file(universe)
# 扫描目录
@@ -383,6 +391,10 @@ class SignalScanner:
if sec_type not in _A_STOCK_TYPES:
continue
# 核心龙头池模式:只保留名单内的代码
if core_codes is not None and code not in core_codes:
continue
market = exchange.upper()
files.append((filepath, market, code))
+13 -4
View File
@@ -199,7 +199,7 @@ class StrengthRanker:
"""扫描全市场并返回强势股排名。
Args:
universe: all/sh/sz/<文件路径>
universe: all/sh/sz/core/<文件路径>core=159 只核心龙头池)
top_n: 返回前 N 名,0=全部
workers: 并发进程数(0=串行,4-8 推荐)
progress_callback: 回调(current, total, name)
@@ -229,13 +229,20 @@ class StrengthRanker:
def _collect_files(self, universe: str) -> list[tuple[Path, str, str]]:
"""收集 A 股 .day 文件列表(复用 scanner 的逻辑)。"""
exchanges: list[str] = []
if universe in ("all", "sz"):
if universe in ("all", "sz", "core"):
exchanges.append("sz")
if universe in ("all", "sh"):
if universe in ("all", "sh", "core"):
exchanges.append("sh")
# 核心龙头池:跨沪深按名单过滤(与 scanner 同一名单)
core_codes: set[str] | None = None
if universe == "core":
from easy_tdx.screen.universe import core_leader_codes
core_codes = core_leader_codes()
# 从文件列表模式读取
if universe not in ("all", "sh", "sz"):
if universe not in ("all", "sh", "sz", "core"):
return self._collect_from_file(universe)
files: list[tuple[Path, str, str]] = []
@@ -247,6 +254,8 @@ class StrengthRanker:
if _detect_security_type(filepath.name) not in _A_STOCK_TYPES:
continue
code = filepath.name.lower()[2:8]
if core_codes is not None and code not in core_codes:
continue
files.append((filepath, exchange.upper(), code))
return files
+190
View File
@@ -0,0 +1,190 @@
"""扫描股票池(universe)定义。
核心龙头池 ``CORE_LEADERS``:159 只,按东方财富全行业龙头名单整理,
涵盖全球第一、国内第一、科技细分龙头与约 40 个行业冠军,四组分层。
来源:社区 Forkswimmingaaron/easy_tdx)按东财名单维护的数据资产,
剥离其个人路径/缓存实现后仅保留静态名单。
接入点:
- :class:`easy_tdx.screen.scanner.SignalScanner` / :class:`easy_tdx.screen.strength.StrengthRanker`
的 ``universe="core"``(离线 .day 扫描按名单过滤,约 3 秒);
- ``GET /api/v1/market/core-leaders``WebUI 展示/导出)。
"""
from __future__ import annotations
__all__ = ["CORE_LEADERS", "CORE_LEADERS_DESC", "core_leader_codes"]
CORE_LEADERS_DESC = "核心龙头池 159 只(东财全行业龙头名单)"
# 分组注释保留名单的层次语义;code → 简称。dict 保序(插入序即展示序)。
CORE_LEADERS: dict[str, str] = {
# 全球第一 / 国际领跑龙头
"002475": "立讯精密",
"002415": "海康威视",
"000725": "京东方A",
"603160": "汇顶科技",
"600745": "闻泰科技",
"002241": "歌尔股份",
"300628": "亿联网络",
"300207": "欣旺达",
"600309": "万华化学",
"300015": "爱尔眼科",
"000661": "长春高新",
"601888": "中国中免",
"601766": "中国中车",
"002050": "三花智控",
"688063": "派能科技",
"688008": "澜起科技",
"600563": "法拉电子",
"688256": "寒武纪",
"600900": "长江电力",
"603993": "洛阳钼业",
"601138": "工业富联",
"300450": "先导智能",
"000338": "潍柴动力",
"601088": "中国神华",
"002714": "牧原股份",
"600660": "福耀玻璃",
"300274": "阳光电源",
"600438": "通威股份",
"002812": "恩捷股份",
"002709": "天赐材料",
"600436": "片仔癀",
"600519": "贵州茅台",
"603288": "海天味业",
"600885": "宏发股份",
"688363": "华熙生物",
"002001": "新和成",
"603260": "合盛硅业",
"600941": "中国移动",
"601728": "中国电信",
# 国内第一 / 行业领军
"000063": "中兴通讯",
"600703": "三安光电",
"600588": "用友网络",
"002230": "科大讯飞",
"601360": "三六零",
"300454": "深信服",
"603019": "中科曙光",
"002410": "广联达",
"002008": "大族激光",
"002371": "北方华创",
"002841": "视源股份",
"300014": "亿纬锂能",
"002439": "启明星辰",
"002916": "深南电路",
"600845": "宝信软件",
"603659": "璞泰来",
"002152": "广电运通",
"300017": "网宿科技",
"000997": "新大陆",
"002396": "星网锐捷",
"002153": "石基信息",
"300271": "华宇软件",
"002405": "四维图新",
"002583": "海能达",
"000050": "深天马A",
"002281": "光迅科技",
"002463": "沪电股份",
# 细分科技与半导体芯片龙头
"300782": "卓胜微",
"300750": "宁德时代",
"000977": "浪潮信息",
"603501": "韦尔股份",
"002938": "鹏鼎控股",
"002600": "领益智造",
"688111": "金山办公",
"300433": "蓝思科技",
"603986": "兆易创新",
"600183": "生益科技",
"300383": "光环新网",
"600536": "中国软件",
"601231": "环旭电子",
"600584": "长电科技",
"300308": "中际旭创",
"603290": "斯达半导",
"300661": "圣邦股份",
"300373": "扬杰科技",
"300666": "江丰电子",
"002236": "大华股份",
"300623": "捷捷微电",
"300349": "金卡智能",
"002079": "苏州固锝",
"603688": "石英股份",
"002119": "康强电子",
"603005": "晶方科技",
"688002": "睿创微纳",
"688099": "晶晨股份",
"002185": "华天科技",
"600460": "士兰微",
"300474": "景嘉微",
"300567": "精测电子",
"300054": "鼎龙股份",
"300398": "飞凯材料",
"300327": "中颖电子",
# 核心大行业与细分龙头
"000858": "五粮液",
"600809": "山西汾酒",
"601100": "恒立液压",
"603638": "艾迪精密",
"000876": "新希望",
"300999": "金龙鱼",
"000895": "双汇发展",
"300059": "东方财富",
"300033": "同花顺",
"600030": "中信证券",
"601318": "中国平安",
"601628": "中国人寿",
"601601": "中国太保",
"000001": "平安银行",
"600036": "招商银行",
"002142": "宁波银行",
"000002": "万科A",
"600048": "保利发展",
"601012": "隆基绿能",
"601636": "旗滨集团",
"002129": "TCL中环",
"300595": "欧普康视",
"600763": "通策医疗",
"000333": "美的集团",
"000651": "格力电器",
"600690": "海尔智家",
"002032": "苏泊尔",
"600887": "伊利股份",
"002460": "赣锋锂业",
"002466": "天齐锂业",
"300122": "智飞生物",
"002007": "华兰生物",
"300142": "沃森生物",
"600276": "恒瑞医药",
"000513": "丽珠集团",
"002271": "东方雨虹",
"002352": "顺丰控股",
"600233": "圆通速递",
"000830": "鲁西化工",
"600426": "华鲁恒升",
"002594": "比亚迪",
"601633": "长城汽车",
"600031": "三一重工",
"000157": "中联重科",
"000425": "徐工机械",
"688981": "中芯国际",
"600547": "山东黄金",
"000975": "山金国际",
"600988": "赤峰黄金",
"600585": "海螺水泥",
"000100": "TCL科技",
"300003": "乐普医疗",
"601668": "中国建筑",
"601390": "中国中铁",
"603799": "华友钴业",
"601899": "紫金矿业",
"002421": "达实智能",
"300223": "北京君正",
}
def core_leader_codes() -> set[str]:
"""核心龙头池 6 位代码集合(scanner 过滤用)。"""
return set(CORE_LEADERS)
+31 -5
View File
@@ -293,6 +293,7 @@ def _create_app(
from easy_tdx.web.routers.finance import router as finance_router
from easy_tdx.web.routers.formula import router as formula_router
from easy_tdx.web.routers.indicator import router as indicator_router
from easy_tdx.web.routers.llm import router as llm_router
from easy_tdx.web.routers.mac_data import router as mac_data_router
from easy_tdx.web.routers.mac_quotes import router as mac_quotes_router
from easy_tdx.web.routers.market import router as market_router
@@ -328,6 +329,8 @@ def _create_app(
app.include_router(strategies_router, prefix="/api/v1")
# 服务器设置路由(列出/测速/切换 TDX host)
app.include_router(server_router, prefix="/api/v1")
# LLM 配置与对话路由(AI 设置页 / AI 解读直连,无行情依赖)
app.include_router(llm_router, prefix="/api/v1")
# 自选股路由(SQLite 持久化,纯 CRUD,不依赖行情连接)
app.include_router(watchlist_router, prefix="/api/v1")
# 实时行情 SSE 路由(依赖 lifespan 里的 QuoteStreamer
@@ -366,19 +369,42 @@ def _create_app(
from starlette.responses import FileResponse
class SPAStaticFiles(StaticFiles):
"""StaticFiles + SPA fallback404 时返回 index.html。"""
"""StaticFiles + SPA fallback404 时返回 index.html。
例外:未匹配的 ``/api/*`` 路径返回 JSON 404 而非 index.html——
SPA fallback 对 API 请求返回 200 HTML 会把"端点不存在/服务是
旧版本"伪装成前端 JSON 解析错误(``Unexpected token '<'``),
且前端 ``resp.ok`` 为 true 连错误分支都不走(v1.29 实测踩坑)。
index.html 一律带 ``Cache-Control: no-store``JS/CSS 是哈希
文件名可以长缓存,但入口 HTML 被缓存会让用户刷新后仍加载旧
资源引用(v1.29 实测:修复已上线、用户强刷仍看到旧版渲染)。
"""
async def get_response(self, path: str, scope): # type: ignore[no-untyped-def]
try:
return await super().get_response(path, scope)
resp = await super().get_response(path, scope)
except Exception:
# 任何 404(路径非文件)都返回 index.html,让前端路由处理。
# 仅对 GET 请求生效;API 路径 (/api/v1/*) 已在前面注册,
# 不会走到这里
# 仅对 GET 请求生效;已注册的 API 路由在路由表命中,不会
# 走到这里——但**未注册**的 /api 路径(如服务加载了旧版
# 本、或端点拼写错)会掉进本 fallback,必须放行 404。
# 注:Windows 下 Starlette 传入的 path 是反斜杠形式,先归一化。
norm_path = path.replace("\\", "/").lstrip("/")
is_api = norm_path == "api" or norm_path.startswith("api/")
if is_api or scope.get("method", "GET") != "GET":
raise
index = _Path(str(self.directory)) / "index.html"
if index.is_file():
return FileResponse(str(index))
return FileResponse(str(index), headers=_INDEX_HEADERS)
raise
# StaticFiles(html=True) 命中目录默认页("/" → index.html)时
# 同样补 no-store,保证入口 HTML 永远取最新
if getattr(resp, "path", "").endswith("index.html"):
resp.headers.update(_INDEX_HEADERS)
return resp
_INDEX_HEADERS = {"Cache-Control": "no-store"}
app.mount("/", SPAStaticFiles(directory=str(dist_dir), html=True), name="web-ui")
logger.info("Web UI mounted from %s (SPA fallback enabled)", dist_dir)
+194
View File
@@ -0,0 +1,194 @@
"""AI 解读历史持久化(Web UI「AI 解读历史」页的数据后端)。
设计对齐 :mod:`easy_tdx.web.watchlist_store`
- 单文件 SQLite,落在统一配置目录(``~/.easy_tdx/llm_history.db``
随 ``EASY_TDX_CONFIG_DIR`` 环境变量走)。
- 短连接 + 写锁串行,跨线程安全(FastAPI 线程池 / task_runner 工作线程内调用)。
- 每次成功的 AI 解读记一条:Prompt(提问上下文)+ 解读正文 + 模型信息 +
策略上下文(策略/参数/标的/周期/日期范围)——策略上下文供历史页
「去回测」一键带参跳转引导。
历史写入属旁路语义:失败不影响解读任务本身(调用方 try/except 兜底)。
"""
from __future__ import annotations
import json
import os
import sqlite3
import threading
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
__all__ = ["LlmHistoryRecord", "LlmHistoryStore", "get_llm_history_store"]
_write_lock = threading.Lock()
def _config_dir() -> Path:
return Path(os.environ.get("EASY_TDX_CONFIG_DIR", str(Path.home() / ".easy_tdx")))
def _now_iso() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
@dataclass
class LlmHistoryRecord:
"""一次成功 AI 解读的完整记录。"""
provider: str
model: str
prompt: str
reply: str
elapsed: float = 0.0
# 策略上下文(「去回测」引导用;手工调用 API 可全部缺省)
strategy: str = ""
strategy_label: str = ""
symbol: str = "" # 6 位代码(与回测页 code 一致)
category: str = ""
params: dict[str, Any] = field(default_factory=dict)
start_date: str = ""
end_date: str = ""
id: int | None = None
created_at: str = ""
def to_dict(self) -> dict[str, Any]:
return {
"id": self.id,
"created_at": self.created_at,
"provider": self.provider,
"model": self.model,
"prompt": self.prompt,
"reply": self.reply,
"elapsed": self.elapsed,
"strategy": self.strategy,
"strategy_label": self.strategy_label,
"symbol": self.symbol,
"category": self.category,
"params": self.params,
"start_date": self.start_date,
"end_date": self.end_date,
}
class LlmHistoryStore:
"""AI 解读历史 SQLite 存储。单例由 :func:`get_llm_history_store` 提供。"""
_SCHEMA = """
CREATE TABLE IF NOT EXISTS llm_history (
id INTEGER PRIMARY KEY AUTOINCREMENT,
created_at TEXT NOT NULL,
provider TEXT NOT NULL DEFAULT '',
model TEXT NOT NULL DEFAULT '',
prompt TEXT NOT NULL DEFAULT '',
reply TEXT NOT NULL DEFAULT '',
elapsed REAL NOT NULL DEFAULT 0,
strategy TEXT NOT NULL DEFAULT '',
strategy_label TEXT NOT NULL DEFAULT '',
symbol TEXT NOT NULL DEFAULT '',
category TEXT NOT NULL DEFAULT '',
params TEXT NOT NULL DEFAULT '{}',
start_date TEXT NOT NULL DEFAULT '',
end_date TEXT NOT NULL DEFAULT ''
);
CREATE INDEX IF NOT EXISTS idx_llm_history_created ON llm_history(created_at DESC);
"""
def __init__(self, db_path: Path | None = None) -> None:
self.db_path = db_path or (_config_dir() / "llm_history.db")
self.db_path.parent.mkdir(parents=True, exist_ok=True)
with self._connect() as conn:
conn.executescript(self._SCHEMA)
def _connect(self) -> sqlite3.Connection:
conn = sqlite3.connect(self.db_path, check_same_thread=False)
conn.row_factory = sqlite3.Row
return conn
def add(self, rec: LlmHistoryRecord) -> LlmHistoryRecord:
"""追加一条记录,返回带 id/created_at 的落库结果。"""
rec.created_at = rec.created_at or _now_iso()
with _write_lock, self._connect() as conn:
cur = conn.execute(
"INSERT INTO llm_history (created_at, provider, model, prompt, reply, elapsed,"
" strategy, strategy_label, symbol, category, params, start_date, end_date)"
" VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
(
rec.created_at,
rec.provider,
rec.model,
rec.prompt,
rec.reply,
rec.elapsed,
rec.strategy,
rec.strategy_label,
rec.symbol,
rec.category,
json.dumps(rec.params, ensure_ascii=False),
rec.start_date,
rec.end_date,
),
)
rec.id = int(cur.lastrowid)
return rec
def list_all(self, limit: int = 50) -> list[LlmHistoryRecord]:
"""按时间倒序列最近 N 条。"""
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM llm_history ORDER BY id DESC LIMIT ?", (int(limit),)
).fetchall()
return [self._row_to_record(r) for r in rows]
def delete(self, record_id: int) -> bool:
"""删除一条;返回是否确实删除。"""
with _write_lock, self._connect() as conn:
cur = conn.execute("DELETE FROM llm_history WHERE id = ?", (int(record_id),))
return cur.rowcount > 0
def clear(self) -> int:
"""清空全部历史;返回删除条数。"""
with _write_lock, self._connect() as conn:
cur = conn.execute("DELETE FROM llm_history")
return cur.rowcount
@staticmethod
def _row_to_record(r: sqlite3.Row) -> LlmHistoryRecord:
try:
params = json.loads(r["params"] or "{}")
except json.JSONDecodeError:
params = {}
return LlmHistoryRecord(
id=int(r["id"]),
created_at=r["created_at"],
provider=r["provider"],
model=r["model"],
prompt=r["prompt"],
reply=r["reply"],
elapsed=float(r["elapsed"] or 0),
strategy=r["strategy"],
strategy_label=r["strategy_label"],
symbol=r["symbol"],
category=r["category"],
params=params if isinstance(params, dict) else {},
start_date=r["start_date"],
end_date=r["end_date"],
)
_store: LlmHistoryStore | None = None
_store_lock = threading.Lock()
def get_llm_history_store() -> LlmHistoryStore:
"""全局单例(首次调用惰性建库)。"""
global _store # noqa: PLW0603 — 模块级单例
if _store is None:
with _store_lock:
if _store is None:
_store = LlmHistoryStore()
return _store
+161 -3
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import logging
from typing import Any
import numpy as np
import pandas as pd
from fastapi import APIRouter, Depends, Query
@@ -26,6 +27,11 @@ router = APIRouter(tags=["bars"])
# 规整后保持的列顺序(匹配旧 SecurityBar 输出契约)
_NORMAL_COLS = ["open", "close", "high", "low", "vol", "amount"]
# 120 分钟线的 category 别名(协议无此枚举,路由层特判)
_MIN_120_ALIASES = frozenset({"MIN_120", "120M", "120MIN"})
# 标准 TdxClient 单次取数上限(60M×2 重采样路径的抓取上限)
_MAX_BARS_PER_FETCH = 800
def _df_resp(df: Any) -> DataFrameResponse:
return DataFrameResponse.from_dataframe(df)
@@ -71,13 +77,151 @@ def _normalize_mac_df(df: pd.DataFrame, daily_plus: bool) -> pd.DataFrame:
return out[cols]
def _resample_pairs(df: pd.DataFrame, count: int) -> pd.DataFrame:
"""相邻两根分钟 bar 聚合成一根(60M×2 → 120M)。
分组规则:从最新端对齐两两配对(奇数根丢最旧一根,保最新数据),
聚合口径 open=first / high=max / low=min / close=last / vol·amount=sum
时间列取配对中后一根。要求 df 按时间升序、含 datetime 列。
Args:
df: 已规整的 60M DataFrame(升序,datetime 列)。
count: 目标 120M 根数(超出的旧数据裁掉)。
Returns:
重采样后的 DataFrame;输入为空时原样返回。
"""
if df is None or df.empty:
return df
out = df.reset_index(drop=True)
if len(out) % 2:
out = out.iloc[1:].reset_index(drop=True) # 丢最旧一根,两两对齐
group = np.arange(len(out)) // 2
agg: dict[str, str] = {"datetime": "last"}
for col, how in (
("open", "first"),
("high", "max"),
("low", "min"),
("close", "last"),
("vol", "sum"),
("amount", "sum"),
):
if col in out.columns:
agg[col] = how
res = out.assign(_g=group).groupby("_g").agg(agg).reset_index(drop=True)
if len(res) > count:
res = res.tail(count).reset_index(drop=True)
return res
def _attach_derived(df: pd.DataFrame) -> pd.DataFrame:
"""每根 bar 附带衍生字段:pre_close / change / change_pct / amplitude_pct。
- ``pre_close``:前一根收盘;首根退化为本根开盘(涨跌记 0)。
- ``change_pct``(close/pre_close - 1)×100。
- ``amplitude_pct``(high - low)/pre_close×100。
- pre_close ≤ 0.01 时按 0.01 兜底(复权后首段价格可能为 0/负,
除零保护;QFQ 负价兜底场景见 /bars 文档)。
"""
if df is None or df.empty or "close" not in df.columns:
return df
out = df.reset_index(drop=True).copy()
close = pd.to_numeric(out["close"], errors="coerce")
pre = close.shift(1)
if "open" in out.columns:
pre = pre.fillna(pd.to_numeric(out["open"], errors="coerce"))
safe_pre = pre.where(pre > 0.01, 0.01)
out["pre_close"] = pre
out["change"] = (close - pre).round(4)
out["change_pct"] = ((close / safe_pre - 1.0) * 100).round(4)
if "high" in out.columns and "low" in out.columns:
high = pd.to_numeric(out["high"], errors="coerce")
low = pd.to_numeric(out["low"], errors="coerce")
out["amplitude_pct"] = ((high - low) / safe_pre * 100).round(4)
return out
async def _fetch_120m(
market: str,
code: str,
start: int,
count: int,
adjust: str,
bar_time: str,
mac_client: Any,
client: Any,
) -> pd.DataFrame:
"""120 分钟 K 线:MAC 原生 times=120 优先,2×60M 重采样兜底。"""
market_value = market_value_from_str(market)
if mac_client is not None:
from easy_tdx.mac.enums import Period
# 1) MAC 原生多分钟线(Period.MINS + times=120
try:
df = await mac_client.get_stock_kline(
market_value,
code,
Period.MINS,
start,
count,
120,
adjust=adjust_from_str(adjust),
bar_time=bar_time,
)
if df is not None and not df.empty:
return _normalize_mac_df(df, daily_plus=False)
_logger.info("/bars MIN_120 原生路径返回空,转 60M 重采样 (%s%s)", market, code)
except Exception as exc: # noqa: BLE001 — 原生不可用时降级,不中断
_logger.warning(
"/bars MIN_120 原生获取失败,转 60M 重采样 (%s%s): %s", market, code, exc
)
# 2) MAC 60M×2 重采样(自动分页,可一次取足 count×2)
try:
df = await mac_client.get_stock_kline(
market_value,
code,
Period.MIN_60,
start,
count * 2,
1,
adjust=adjust_from_str(adjust),
bar_time=bar_time,
)
res = _resample_pairs(_normalize_mac_df(df, daily_plus=False), count)
if res is not None and not res.empty:
return res
except Exception as exc: # noqa: BLE001
_logger.warning("/bars MIN_120 60M重采样(MAC)失败 (%s%s): %s", market, code, exc)
# 3) 标准 TdxClient 60M×2(无 MAC;单次上限 800 根 → 最多 400 根 120M
fetch_n = min(count * 2, _MAX_BARS_PER_FETCH)
if fetch_n < count * 2:
_logger.info(
"/bars MIN_120 回退路径单次上限 %d 根 60M,最多合成 %d 根 120M",
_MAX_BARS_PER_FETCH,
_MAX_BARS_PER_FETCH // 2,
)
df = await client.get_security_bars(
market_from_str(market), code, category_from_str("MIN_60"), start, fetch_n,
bar_time=bar_time,
)
return _resample_pairs(df, count)
@router.get("/bars", response_model=DataFrameResponse)
async def security_bars(
market: str = Query(..., description="市场: SZ, SH, BJ"),
code: str = Query(..., min_length=6, max_length=6),
category: str = Query(
"DAY",
description="K线周期: MIN_1, MIN_5, MIN_15, MIN_30, MIN_60, DAY, WEEK, MONTH, YEAR",
description=(
"K线周期: MIN_1, MIN_5, MIN_15, MIN_30, MIN_60, MIN_120(120分钟), "
"DAY, WEEK, MONTH, SEASON, YEAR"
),
),
start: int = Query(0, ge=0),
count: int = Query(800, ge=1, le=800),
@@ -96,9 +240,21 @@ async def security_bars(
MAC 主机未连接时自动回退 AsyncTdxClient.get_security_bars(无复权,adjust 参数忽略)。
输出契约与旧版一致:日线返回 ``date`` 列,分钟线返回 ``datetime`` 列。
``category=MIN_120`` 为 120 分钟线:MAC 原生 ``Period.MINS × times=120``
优先,失败则取 2 倍 60M 数据相邻两根聚合(open=first/high=max/low=min/
close=last/vol·amount=sum),标准客户端回退路径最多合成 400 根。
每根 bar 附带衍生字段:``pre_close``(前收,首根=本根开盘)、``change``、
``change_pct``、``amplitude_pct``(振幅%)。pre_close ≤ 0.01 时按 0.01
兜底(QFQ 复权后早期价格可能为 0/负)。
vol 单位:分钟线/日线 = 成交量(股);周/月/季/年线服务端原样返回真实
成交量/100,回退路径(标准 TdxClient)已 ×100 还原为股。
"""
if category.upper() in _MIN_120_ALIASES:
df = await _fetch_120m(market, code, start, count, adjust, bar_time, mac_client, client)
return _df_resp(_attach_derived(df))
cat = category_from_str(category)
if mac_client is not None:
period, times = period_times_from_category(cat)
@@ -123,7 +279,7 @@ async def security_bars(
df = await client.get_security_bars(
market_from_str(market), code, cat, start, count, bar_time=bar_time
)
return _df_resp(df)
return _df_resp(_attach_derived(df))
@router.get("/bars/index", response_model=DataFrameResponse)
@@ -143,11 +299,13 @@ async def index_bars(
vol 单位:日线/周线/月线/季线/年线 = 成交量(手)(周及以上周期服务端
原样返回真实成交量/100,已 ×100 还原);**分钟线协议不提供成交量**
(报文中该字段实为成交额/100),vol 为 ``null``,请勿当作成交量使用。
每根 bar 同样附带 ``pre_close/change/change_pct/amplitude_pct`` 衍生字段。
"""
df = await client.get_index_bars(
market_from_str(market), code, category_from_str(category), start, count, bar_time=bar_time
)
return _df_resp(df)
return _df_resp(_attach_derived(df))
@router.get("/minute", response_model=DataFrameResponse)
+273
View File
@@ -0,0 +1,273 @@
"""LLM 配置与对话路由(WebUI「AI 设置」页 + AI 解读直连)。"""
from __future__ import annotations
import asyncio
import logging
import time
from typing import Any
from fastapi import APIRouter
from pydantic import BaseModel, Field
from easy_tdx.ai.llm import (
PROVIDER_PRESETS,
LlmClient,
LlmConfig,
LlmError,
config_path,
load_config,
mask_key,
resolve_config,
save_config,
)
from easy_tdx.web.backtest_schemas import TaskStateResponse, TaskSubmitResponse
from easy_tdx.web.task_runner import get_runner
logger = logging.getLogger(__name__)
router = APIRouter(tags=["llm"])
class LlmConfigUpdate(BaseModel):
"""PUT /llm/config 请求体。
``api_key`` 缺省或等于当前脱敏回显值时保留原 key——前端把脱敏串原样
回传不会把真 key 冲掉;只有填了新值才覆盖。
"""
provider: str = Field("deepseek", description="Provider id(见 GET /llm/config providers")
api_url: str = Field("", description="API 地址,空 = 用该 Provider 预设")
api_key: str = Field("", description="API Key(留空/传回脱敏串 = 不修改已存 key)")
model: str = Field("", description="模型名,空 = 用该 Provider 默认模型")
temperature: float = Field(0.3, ge=0.0, le=2.0)
max_tokens: int = Field(
16000, ge=64, le=128_000, description="输出上限;思考型模型的思考链计入此预算,建议 ≥16000"
)
timeout: float = Field(180.0, ge=5.0, le=600.0, description="读超时;报告解读建议 ≥120")
system_prompt: str = ""
class LlmChatContext(BaseModel):
"""AI 解读附带的策略上下文(历史页「去回测」引导用,全部可缺省)。"""
strategy: str = Field("", description="策略注册表 key(如 ma_cross")
strategy_label: str = Field("", description="策略中文名")
symbol: str = Field("", description="6 位标的代码")
category: str = Field("", description="K 线周期")
params: dict[str, Any] = Field(default_factory=dict, description="策略参数")
start_date: str = ""
end_date: str = ""
class LlmChatRequest(BaseModel):
"""POST /llm/chat(/async) 请求体(如把 AI 解读 Prompt 直接发给已配置的 LLM)。"""
prompt: str = Field(..., min_length=1, max_length=200_000)
system_prompt: str | None = Field(None, description="None = 用配置里的默认系统提示")
override: LlmConfigUpdate | None = Field(
None, description="临时覆盖配置(不落盘,仅本次调用)"
)
context: LlmChatContext | None = Field(
None, description="策略上下文(随成功解读一并落历史库)"
)
def _merge_api_key(submitted: str, current: str) -> str:
"""按表单语义合并 api_key:留空/回传脱敏串 = 沿用;CLEAR = 清除;其余 = 覆盖。
前端把脱敏串原样回传不会把真 key 冲掉;显式填 ``CLEAR`` 可移除已存
key(否则换 Provider 时旧 key 会残留且 UI 无清除入口)。
"""
key = submitted.strip()
if not key or key == mask_key(current):
return current
if key.upper() == "CLEAR":
return ""
return key
def _override_config(override: LlmConfigUpdate | None) -> LlmConfig | None:
"""请求体临时配置 → LlmConfig(不落盘)。api_key 走 _merge_api_key 语义。"""
if override is None:
return None
current = load_config()
key = _merge_api_key(override.api_key, current.api_key)
return LlmConfig(
provider=override.provider,
api_url=override.api_url.strip(),
api_key=key,
model=override.model.strip(),
temperature=override.temperature,
max_tokens=override.max_tokens,
timeout=override.timeout,
)
@router.get("/llm/config")
async def get_llm_config() -> dict[str, Any]:
"""当前 LLM 配置(key 脱敏)+ Provider 预设表 + 配置文件路径。"""
cfg = load_config()
preset = PROVIDER_PRESETS.get(cfg.provider, PROVIDER_PRESETS["custom"])
try:
resolved = resolve_config(cfg)
missing: list[str] = []
if preset.needs_key and not cfg.api_key:
missing.append("api_key")
except ValueError as exc:
resolved = cfg # type: ignore[assignment]
missing = [str(exc)]
return {
"config": cfg.to_dict(mask_api_key=True),
"providers": [p.to_dict() for p in PROVIDER_PRESETS.values()],
"configured": not missing,
"missing": missing,
"config_path": str(config_path()),
"resolved": {"api_url": resolved.api_url, "model": resolved.model},
}
@router.put("/llm/config")
async def update_llm_config(req: LlmConfigUpdate) -> dict[str, Any]:
"""保存 LLM 配置到 llm.json(WebUI 与手工编辑同一份文件,双向兼容)。"""
if req.provider not in PROVIDER_PRESETS:
valid = ", ".join(PROVIDER_PRESETS)
raise ValueError(f"未知 provider '{req.provider}',可选: {valid}")
current = load_config()
new_key = _merge_api_key(req.api_key, current.api_key)
cfg = LlmConfig(
provider=req.provider,
api_url=req.api_url.strip(),
api_key=new_key,
model=req.model.strip(),
temperature=req.temperature,
max_tokens=req.max_tokens,
timeout=req.timeout,
system_prompt=req.system_prompt or current.system_prompt,
)
path = save_config(cfg)
logger.info("LLM 配置已保存: provider=%s model=%s (%s)", cfg.provider, cfg.model, path)
return {"ok": True, "config_path": str(path), "config": cfg.to_dict(mask_api_key=True)}
def _record_history(
provider: str, model: str, prompt: str, reply: str, elapsed: float,
ctx: LlmChatContext | None,
) -> None:
"""成功解读旁路落库(llm_history.db)。失败只记日志,不影响解读结果。"""
from easy_tdx.web.llm_history_store import LlmHistoryRecord, get_llm_history_store
try:
get_llm_history_store().add(
LlmHistoryRecord(
provider=provider,
model=model,
prompt=prompt,
reply=reply,
elapsed=elapsed,
**(ctx.model_dump() if ctx else {}),
)
)
except Exception: # noqa: BLE001 — 历史属旁路语义
logger.exception("AI 解读历史落库失败(不影响解读结果)")
@router.post("/llm/test")
async def test_llm(override: LlmConfigUpdate | None = None) -> dict[str, Any]:
"""连通性测试:用已保存配置(或请求体内临时配置)发一句极短 ping。"""
return await LlmClient(_override_config(override)).test()
@router.post("/llm/chat")
async def llm_chat(req: LlmChatRequest) -> dict[str, Any]:
"""一轮 LLM 对话:把 prompt(如回测报告解读 Prompt)发给已配置的模型。"""
client = LlmClient(_override_config(req.override))
try:
t0 = time.perf_counter()
reply = await client.chat(req.prompt, system_prompt=req.system_prompt)
except LlmError as exc:
# 全局 ValueError 处理器 → 400 {error, detail},前端 formatError 可读展示
raise ValueError(str(exc)) from exc
elapsed = round(time.perf_counter() - t0, 1)
_record_history(
client.config.provider, client.config.model, req.prompt, reply, elapsed, req.context
)
return {"reply": reply, "model": client.config.model, "provider": client.config.provider}
@router.post("/llm/chat/async", response_model=TaskSubmitResponse, status_code=202)
async def llm_chat_async(req: LlmChatRequest) -> TaskSubmitResponse:
"""提交 AI 解读后台任务(长耗时模型调用不占住 HTTP 连接)。
大报告解读 1-3 分钟,同步 HTTP 等待对代理/浏览器都不友好;这里接入
与回测同一套任务执行器(``task_runner``4 线程池 + SQLite 持久化),
前端短轮询 ``GET /llm/chat/tasks/{task_id}`` 取状态,断线重连后仍可
查询。配置不完整(缺 url/model)在提交时即报 400;网络/鉴权/超时
类错误发生在任务内,体现在 TaskState.error。
"""
client = LlmClient(_override_config(req.override)) # 提交期即校验配置
desc = f"AI 解读 | {client.config.provider} · {client.config.model} | {len(req.prompt)}"
def _run() -> dict[str, Any]:
t0 = time.perf_counter()
reply = asyncio.run(client.chat(req.prompt, system_prompt=req.system_prompt))
elapsed = round(time.perf_counter() - t0, 1)
_record_history(
client.config.provider, client.config.model, req.prompt, reply, elapsed, req.context
)
return {"reply": reply, "model": client.config.model, "provider": client.config.provider,
"elapsed": elapsed}
runner = get_runner()
task_id = runner.submit(_run, description=desc)
state = runner.get(task_id)
status: Any = state.status if state.status in ("pending", "running") else "running"
return TaskSubmitResponse(task_id=task_id, status=status)
@router.get("/llm/chat/tasks/{task_id}", response_model=TaskStateResponse)
async def llm_chat_task(task_id: str) -> TaskStateResponse:
"""查询 AI 解读任务状态(与回测任务同一存储,语义化路径别名)。"""
runner = get_runner()
try:
state = runner.get(task_id)
except KeyError as exc:
raise ValueError(str(exc)) from exc
return TaskStateResponse(
task_id=state.task_id,
status=state.status,
result=state.result,
error=state.error,
description=state.description,
elapsed=(state.finished_at or time.time()) - (state.started_at or state.created_at),
)
@router.get("/llm/history")
async def list_llm_history(limit: int = 50) -> dict[str, Any]:
"""AI 解读历史(时间倒序)。每条含 Prompt、解读正文与策略上下文。"""
from easy_tdx.web.llm_history_store import get_llm_history_store
items = get_llm_history_store().list_all(limit=min(max(limit, 1), 200))
return {"items": [r.to_dict() for r in items], "count": len(items)}
@router.delete("/llm/history/{record_id}")
async def delete_llm_history(record_id: int) -> dict[str, Any]:
"""删除一条历史记录。"""
from easy_tdx.web.llm_history_store import get_llm_history_store
if not get_llm_history_store().delete(record_id):
raise ValueError(f"历史记录 {record_id} 不存在")
return {"ok": True}
@router.delete("/llm/history")
async def clear_llm_history() -> dict[str, Any]:
"""清空全部历史记录。"""
from easy_tdx.web.llm_history_store import get_llm_history_store
deleted = get_llm_history_store().clear()
return {"ok": True, "deleted": deleted}
+29 -1
View File
@@ -76,6 +76,34 @@ async def market_stat(
return _df_response(df)
@router.get("/market/session")
async def market_session() -> dict[str, Any]:
"""A 股有效行情时段判断(供前端自动刷新门控校准)。
窗口 09:15~11:30、13:00~15:05(含集合竞价缓冲,午休除外),周一至周五。
节假日不做日历判断——盘外误判为盘中只会多拉一次快照,无副作用。
"""
from easy_tdx.realtime.session import session_info
return session_info()
@router.get("/market/core-leaders", response_model=DataFrameResponse)
async def core_leaders() -> DataFrameResponse:
"""核心龙头池(159 只,按东方财富全行业龙头名单整理)。
数据资产供前端展示/导出;扫描场景走 ``universe="core"``screen scan
与 /market/strength 均支持)。
"""
from easy_tdx.screen.universe import CORE_LEADERS
rows = [
{"code": code, "name": name, "market": "SH" if code.startswith(("6", "9")) else "SZ"}
for code, name in CORE_LEADERS.items()
]
return DataFrameResponse(data=rows, count=len(rows))
@router.get("/fund-flow", response_model=DataFrameResponse)
async def fund_flow(
market: str = Query(..., description="市场: SZ, SH"),
@@ -117,7 +145,7 @@ async def market_strength(
w60: float | None = Query(None, description="自定义 60 日权重(覆盖预设)"),
vol_adjusted: bool | None = Query(None, description="波动率惩罚开关(覆盖预设)"),
top_n: int = Query(50, ge=1, le=5000, description="返回前 N 名"),
universe: str = Query("all", description="范围: all/sh/sz"),
universe: str = Query("all", description="范围: all/sh/sz/corecore=核心龙头池159只)"),
min_listed_days: int = Query(65, ge=30, description="最小上市天数"),
min_amount: float = Query(0.0, ge=0, description="最近 5 日日均成交额下限(元)"),
vipdoc: str | None = Query(None, description="离线数据目录(默认自动检测)"),