Files
xiaobaifupan/next/backend/features/screener/compiler.py
T

114 lines
4.3 KiB
Python
Raw 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.
from __future__ import annotations
import hashlib
import json
from typing import Any
from backend.features.screener.catalog import CatalogError, factor_catalog, validate_formula
PROMPT_VERSION = "screener:formula-compiler:v1"
def compile_messages(description: str) -> list[dict[str, str]]:
factors = factor_catalog()["factors"]
schema = {
"universe": {"exclude_st": True, "listed_days_min": 120},
"filters": [{"field": "amount_billion", "op": ">=", "value": 1}],
"score": [{"field": "return_20d", "weight": 1.0, "direction": "desc"}],
"limit": 30,
"min_score": 0.5,
}
return [
{
"role": "system",
"content": (
"你是受控选股公式编译器。只输出一个JSON对象,不得输出Markdown、解释、股票或代码。"
"只能使用给定因子;filters运算符仅限 >、>=、<、<=、==、!=、between、in"
"score权重必须为0至1小数且总和等于1direction仅限asc或desc"
"limit为1至50min_score为0至1。无法完全表达时选择最接近的已知因子,不得创造字段。"
f"\nJSON结构:{json.dumps(schema, ensure_ascii=False, separators=(',', ':'))}"
f"\n可用因子:{json.dumps(factors, ensure_ascii=False, separators=(',', ':'))}"
),
},
{"role": "user", "content": description},
]
def description_key(description: str) -> str:
digest = hashlib.sha256(description.encode("utf-8")).hexdigest()[:16]
return f"formula:{digest}"
def parse_compiled_formula(content: str) -> dict[str, Any]:
text = content.strip()
if text.startswith("```"):
lines = text.splitlines()
if len(lines) < 3 or lines[-1].strip() != "```":
raise CatalogError("模型返回的公式格式无效")
text = "\n".join(lines[1:-1]).strip()
try:
raw = json.loads(text)
except json.JSONDecodeError as exc:
raise CatalogError("模型未返回有效JSON公式") from exc
if not isinstance(raw, dict):
raise CatalogError("模型返回的公式必须是对象")
return normalize_custom_formula(raw)
def normalize_custom_formula(raw: dict[str, Any]) -> dict[str, Any]:
universe = raw.get("universe")
filters = raw.get("filters")
scores = raw.get("score")
if not isinstance(universe, dict) or not isinstance(filters, list) or not isinstance(
scores, list
):
raise CatalogError("模型返回的公式结构不完整")
exclude_st = universe.get("exclude_st", True)
listed_days = universe.get("listed_days_min", 120)
if not isinstance(exclude_st, bool) or not isinstance(listed_days, int):
raise CatalogError("股票范围设置无效")
normalized_scores = []
for score in scores:
if not isinstance(score, dict):
raise CatalogError("评分因子结构无效")
normalized_scores.append(
{
"field": score.get("field"),
"weight": score.get("weight"),
"direction": score.get("direction", "desc"),
}
)
if any(not isinstance(item, dict) for item in filters):
raise CatalogError("筛选条件结构无效")
numeric_weights = [item["weight"] for item in normalized_scores]
if all(
isinstance(value, (int, float)) and not isinstance(value, bool)
for value in numeric_weights
):
total = sum(float(value) for value in numeric_weights)
if abs(total - 100) <= 0.0001:
for item in normalized_scores:
item["weight"] = float(item["weight"]) / 100
minimum = raw.get("min_score")
if isinstance(minimum, (int, float)) and not isinstance(minimum, bool) and 1 < minimum <= 100:
minimum = float(minimum) / 100
formula = {
"universe": {
"exclude_st": exclude_st,
"listed_days_min": listed_days,
},
"filters": [
{
"field": item.get("field"),
"op": item.get("op"),
"value": item.get("value"),
}
for item in filters
],
"score": normalized_scores,
"limit": raw.get("limit"),
"min_score": minimum,
}
validate_formula(formula)
return formula