Files
stock/backend/app/screener/llm.py
2026-08-15 08:57:15 +08:00

218 lines
13 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""LLM 条件解析器DeepSeekOpenAI 兼容 /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_jKDJ 的 K/D/J 值、rsi、macd_dif / macd_dea / macd_histMACD 的 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_deavalue_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): # 未带版本段则补 /v1DeepSeek/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 Keyplatform.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_jKDJ 的 K/D/J 值、rsi、macd_dif / macd_dea / macd_histMACD 的 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 Keyplatform.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}")