218 lines
13 KiB
Python
218 lines
13 KiB
Python
"""LLM 条件解析器(DeepSeek,OpenAI 兼容 /chat/completions)。
|
||
|
||
把自然语言选股需求解析成 ScreenConditions(结构化 JSON)。
|
||
- response_format=json_object + temperature=0.1 保证结构稳定
|
||
- 解析/校验失败带错误重试 1 次
|
||
- 上游错误信息透传给前端(ScreenerError)
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
import re
|
||
|
||
import httpx
|
||
|
||
from ..config import settings
|
||
from ..schemas import EventBacktestSpec, ScreenConditions
|
||
|
||
SYSTEM_PROMPT = """你是 A 股选股条件解析器。把用户的自然语言解析成一个 JSON 对象,只输出 JSON,不要任何解释、注释或代码块围栏。完全无法理解时输出 {"error": "原因"}。
|
||
|
||
输出结构:
|
||
{"indicator": [...], "snapshot": [...], "exclude_st": true, "exclude_delisted": true, "exclude_bj": true}
|
||
(indicator 与 snapshot 至少一个非空;用户没有提到的条件不要编造)
|
||
|
||
【indicator 数组】技术指标条件,元素字段:
|
||
- "indicator": 指标名,白名单:kdj_k / kdj_d / kdj_j(KDJ 的 K/D/J 值)、rsi、macd_dif / macd_dea / macd_hist(MACD 的 DIF/DEA/柱)、ma(收盘价均线)、boll_upper / boll_mid / boll_lower(布林轨道)、close(收盘价)、pct_chg(日涨跌幅%)
|
||
- "params": 指标参数(可选),默认:KDJ {"n":9,"m1":3,"m2":3};RSI {"period":14};MACD {"fast":12,"slow":26,"signal":9};MA {"period":20};BOLL {"period":20,"std":2}
|
||
- "op": "gt" | "ge" | "lt" | "le" | "between"
|
||
- "value": 比较数值(between 时为下界),"value2": between 上界
|
||
- "value_indicator": 可选。指标与指标比较时填另一指标名(同白名单),如 "DIF大于DEA" -> indicator=macd_dif, op=gt, value_indicator=macd_dea, value=0;"股价在布林带下轨之下" -> indicator=close, op=lt, value_indicator=boll_lower, value=0
|
||
- "value_params": 可选。比较对象指标需要不同参数时指定,如 "MA5 上穿 MA20" -> indicator=ma, params={"period":5}, op=gt, value_indicator=ma, value_params={"period":20}, value=0;不填则比较对象沿用 params 中适用于它的参数或默认参数
|
||
- "lookback": 检查最近 N 个交易日(默认 1)
|
||
- "match": "all"(窗口内每天满足,默认)或 "any"(窗口内任一天满足)
|
||
|
||
【snapshot 数组】最新交易日截面条件,元素字段:
|
||
- "field": 白名单:total_mv(总市值)、circ_mv(流通市值)、pe_ttm(市盈率TTM)、pb(市净率)、turnover_rate(换手率)、close(最新价)
|
||
- "op"/"value"/"value2" 同上
|
||
|
||
【单位约定】市值条件统一用亿元(如"市值大于100亿,小于200亿"→ between 100~200);换手率用百分数值("换手率大于5%"→ value 5);价格类用元;pe/pb 用倍数。
|
||
|
||
【时间语义】"今天/今日"→ lookback=1;"这两天/最近N天/连续N日"→ lookback=N 且 match="all";"近N日内曾经/任一天"→ lookback=N 且 match="any"。只支持以最新交易日为终点的窗口,不要生成具体某一天的条件。
|
||
|
||
【排除规则】默认 exclude_st=true、exclude_delisted=true、exclude_bj=true;用户明确说"包含北交所/包含ST"时才把对应项设为 false。
|
||
|
||
【交叉类表述的近似】"MACD金叉/刚金叉"用 macd_dif gt macd_dea(value_indicator)+ 适当 lookback/match 近似;"跌破均线"用 close lt ma 近似;无法近似表达的复杂条件直接忽略,保留可表达的部分。
|
||
|
||
示例1:
|
||
输入:帮我找出这两天 KDJ 中的 J 小于 10,市值大于 100 亿,小于 200 亿的公司
|
||
输出:{"indicator":[{"indicator":"kdj_j","params":{"n":9,"m1":3,"m2":3},"op":"lt","value":10,"lookback":2,"match":"all"}],"snapshot":[{"field":"total_mv","op":"between","value":100,"value2":200}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true}
|
||
|
||
示例2:
|
||
输入:RSI 低于 30,市盈率 TTM 小于 20 的公司
|
||
输出:{"indicator":[{"indicator":"rsi","params":{"period":14},"op":"lt","value":30,"lookback":1,"match":"all"}],"snapshot":[{"field":"pe_ttm","op":"lt","value":20}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true}
|
||
|
||
示例3:
|
||
输入:近 5 天曾经 MACD 金叉(DIF 上穿 DEA),换手率大于 5%,流通市值小于 100 亿
|
||
输出:{"indicator":[{"indicator":"macd_dif","params":{"fast":12,"slow":26,"signal":9},"op":"gt","value":0,"value_indicator":"macd_dea","lookback":5,"match":"any"}],"snapshot":[{"field":"turnover_rate","op":"gt","value":5},{"field":"circ_mv","op":"lt","value":100}],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true}"""
|
||
|
||
|
||
class ScreenerError(RuntimeError):
|
||
"""选股链路可预期的业务错误(信息可直接透传给前端)。"""
|
||
|
||
|
||
def _endpoint(base_url: str) -> str:
|
||
"""归一化 base_url -> 完整 chat/completions URL。
|
||
|
||
兼容多种写法:DeepSeek/OpenAI 的 .../v1、智谱 GLM 的 .../v4、或直接给完整路径。
|
||
"""
|
||
base = base_url.rstrip("/")
|
||
if base.endswith("/chat/completions"):
|
||
return base
|
||
if not re.search(r"/v\d+$", base): # 未带版本段则补 /v1(DeepSeek/OpenAI 惯例)
|
||
base += "/v1"
|
||
return f"{base}/chat/completions"
|
||
|
||
|
||
def _extract_json(content: str) -> dict:
|
||
"""从 LLM 输出提取 JSON:剥代码围栏,或取首 { 到末 } 的子串。"""
|
||
text = content.strip()
|
||
if text.startswith("```"):
|
||
# 剥 ```json ... ``` 围栏
|
||
text = text.split("```", 2)[1]
|
||
if text.startswith("json"):
|
||
text = text[4:]
|
||
text = text.strip()
|
||
if not text.startswith("{"):
|
||
start, end = text.find("{"), text.rfind("}")
|
||
if start < 0 or end <= start:
|
||
raise ValueError("输出中不含 JSON 对象")
|
||
text = text[start : end + 1]
|
||
obj = json.loads(text)
|
||
if not isinstance(obj, dict):
|
||
raise ValueError("JSON 不是对象")
|
||
return obj
|
||
|
||
|
||
def _build_messages(text: str, retry_error: str | None = None) -> list[dict]:
|
||
"""构造 system + user 消息;retry 时附上一次解析错误要求修正。"""
|
||
user = f"解析以下选股需求:{text}"
|
||
if retry_error:
|
||
user += f"\n\n上一次输出无法通过校验,错误:{retry_error}。请修正后重新只输出 JSON。"
|
||
return [{"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user}]
|
||
|
||
|
||
async def _chat(messages: list[dict]) -> str:
|
||
"""调 OpenAI 兼容接口,返回 assistant 文本。上游错误抛 ScreenerError。"""
|
||
body = {
|
||
"model": settings.llm_model,
|
||
"messages": messages,
|
||
"temperature": 0.1,
|
||
"response_format": {"type": "json_object"},
|
||
"max_tokens": 2000,
|
||
}
|
||
headers = {"Authorization": f"Bearer {settings.llm_api_key}"}
|
||
async with httpx.AsyncClient(timeout=settings.llm_timeout) as client:
|
||
try:
|
||
r = await client.post(_endpoint(settings.llm_base_url), json=body, headers=headers)
|
||
except httpx.HTTPError as e: # 网络/超时
|
||
raise ScreenerError(f"LLM 服务无法访问({settings.llm_base_url}): {e}") from e
|
||
if r.status_code >= 400:
|
||
detail = ""
|
||
try:
|
||
detail = r.json().get("error", {}).get("message", "")
|
||
except Exception: # noqa: BLE001
|
||
detail = r.text[:200]
|
||
raise ScreenerError(f"LLM 接口错误 (HTTP {r.status_code}): {detail or '无详细信息'}")
|
||
try:
|
||
return r.json()["choices"][0]["message"]["content"] or ""
|
||
except (KeyError, IndexError, TypeError) as e:
|
||
raise ScreenerError(f"LLM 返回结构异常: {r.text[:200]}") from e
|
||
|
||
|
||
async def parse_conditions(text: str) -> ScreenConditions:
|
||
"""主入口:自然语言 -> ScreenConditions。未配 key / 解析两次失败抛 ScreenerError。"""
|
||
if not settings.llm_api_key:
|
||
raise ScreenerError(
|
||
"未配置 LLM_API_KEY:请在 backend/.env 填入 DeepSeek API Key(platform.deepseek.com 获取)后重启后端"
|
||
)
|
||
|
||
retry_error: str | None = None
|
||
for _ in range(2): # 首次 + 失败重试 1 次
|
||
content = await _chat(_build_messages(text, retry_error))
|
||
try:
|
||
obj = _extract_json(content)
|
||
if "error" in obj and not obj.get("indicator") and not obj.get("snapshot"):
|
||
raise ScreenerError(f"AI 无法理解该选股需求:{obj['error']}")
|
||
return ScreenConditions.model_validate(obj)
|
||
except ScreenerError:
|
||
raise
|
||
except Exception as e: # noqa: BLE001 —— JSON/校验失败,带错误重试
|
||
retry_error = str(e)[:300]
|
||
raise ScreenerError(f"AI 解析结果两次未通过校验,最后错误:{retry_error}")
|
||
|
||
|
||
# ---------- 事件回测解析(自然语言 -> EventBacktestSpec) ----------
|
||
|
||
EVENT_SYSTEM_PROMPT = """你是 A 股事件回测参数解析器。用户描述一个「入场信号 + 买卖时机 + 持有期」的事件回测需求,把它解析成 JSON,只输出 JSON,不要任何解释或代码块围栏。完全无法理解时输出 {"error": "原因"}。
|
||
|
||
输出结构:
|
||
{"entry": {"indicator": [...], "snapshot": [], "exclude_st": true, "exclude_delisted": true, "exclude_bj": true}, "entry_timing": "next_open", "holding_days": 3, "exit_timing": "close"}
|
||
|
||
【entry.indicator 数组】入场信号条件(必填,至少 1 条),元素字段与白名单:
|
||
- "indicator": kdj_k / kdj_d / kdj_j(KDJ 的 K/D/J 值)、rsi、macd_dif / macd_dea / macd_hist(MACD 的 DIF/DEA/柱)、ma(收盘价均线)、boll_upper / boll_mid / boll_lower(布林轨道)、close(收盘价)、pct_chg(日涨跌幅%)
|
||
- "params": 指标参数(可选),默认:KDJ {"n":9,"m1":3,"m2":3};RSI {"period":14};MACD {"fast":12,"slow":26,"signal":9};MA {"period":20};BOLL {"period":20,"std":2}
|
||
- "op": "gt" | "ge" | "lt" | "le" | "between";"value"(between 时为下界)、"value2"(上界)
|
||
- "value_indicator": 指标与指标比较时填另一指标名(同白名单),如 "DIF 大于 DEA" -> indicator=macd_dif, op=gt, value_indicator=macd_dea, value=0
|
||
- "value_params": 比较对象指标参数不同时指定,如 "MA5 上穿 MA20" -> indicator=ma, params={"period":5}, op=gt, value_indicator=ma, value_params={"period":20}, value=0
|
||
- "lookback": 信号需连续/曾经满足的交易日窗口(默认 1)
|
||
- "match": "all"(窗口内每天满足,默认)或 "any"(窗口内任一天满足)
|
||
|
||
【时间语义】"连续三天 J 小于 10" -> lookback=3, match="all";"近 5 天曾经金叉" -> lookback=5, match="any"。
|
||
|
||
【entry_timing】买入时机:"第二天开盘购买/次日开盘买入" -> "next_open"(默认);"第二天收盘买入" -> "next_close"。
|
||
|
||
【holding_days】买入后持有 N 个交易日(int,默认 3)。"未来三天的涨幅" -> holding_days=3;"持有 10 天" -> 10;"持有一个月" -> 20。
|
||
|
||
【exit_timing】到期卖出价:"close"(收盘卖,默认)或 "open"(开盘卖)。
|
||
|
||
【entry.snapshot】截面过滤条件一般不适用于历史回测,除非用户明确说"只回测市值大于 X 亿的股票"才填,其余情况留空数组。
|
||
|
||
示例:
|
||
输入:在连续三天 J 小于 10 的时候第二天开盘购买,之后未来三天的涨幅有多少
|
||
输出:{"entry":{"indicator":[{"indicator":"kdj_j","params":{"n":9,"m1":3,"m2":3},"op":"lt","value":10,"lookback":3,"match":"all"}],"snapshot":[],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true},"entry_timing":"next_open","holding_days":3,"exit_timing":"close"}
|
||
|
||
示例:
|
||
输入:RSI 低于 30 的第二天开盘买入持有 5 天收盘卖出
|
||
输出:{"entry":{"indicator":[{"indicator":"rsi","params":{"period":14},"op":"lt","value":30,"lookback":1,"match":"all"}],"snapshot":[],"exclude_st":true,"exclude_delisted":true,"exclude_bj":true},"entry_timing":"next_open","holding_days":5,"exit_timing":"close"}"""
|
||
|
||
|
||
def _build_event_messages(text: str, retry_error: str | None = None) -> list[dict]:
|
||
user = f"解析以下事件回测需求:{text}"
|
||
if retry_error:
|
||
user += f"\n\n上一次输出无法通过校验,错误:{retry_error}。请修正后重新只输出 JSON。"
|
||
return [{"role": "system", "content": EVENT_SYSTEM_PROMPT}, {"role": "user", "content": user}]
|
||
|
||
|
||
async def parse_event_spec(text: str) -> EventBacktestSpec:
|
||
"""自然语言 -> EventBacktestSpec。复用 _chat/_extract_json,失败带错误重试 1 次。"""
|
||
if not settings.llm_api_key:
|
||
raise ScreenerError(
|
||
"未配置 LLM_API_KEY:请在 backend/.env 填入 DeepSeek API Key(platform.deepseek.com 获取)后重启后端"
|
||
)
|
||
retry_error: str | None = None
|
||
for _ in range(2):
|
||
content = await _chat(_build_event_messages(text, retry_error))
|
||
try:
|
||
obj = _extract_json(content)
|
||
if "error" in obj and not obj.get("entry"):
|
||
raise ScreenerError(f"AI 无法理解该回测需求:{obj['error']}")
|
||
spec = EventBacktestSpec.model_validate(obj)
|
||
if not spec.entry.indicator:
|
||
raise ValueError("entry.indicator 不能为空")
|
||
return spec
|
||
except ScreenerError:
|
||
raise
|
||
except Exception as e: # noqa: BLE001
|
||
retry_error = str(e)[:300]
|
||
raise ScreenerError(f"AI 解析回测参数两次未通过校验,最后错误:{retry_error}")
|