Files
Mutual_Fund/nl2sql/query_experience.py
T

78 lines
2.6 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""NL2SQL 查询体验辅助能力。"""
from __future__ import annotations
import csv
import re
from io import StringIO
from sqlglot import exp, parse_one
from nl2sql.semantics import resolve_semantics
_IDENTIFIER = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
_TIME_MARKERS = ("近", "本年", "今年", "去年", "上月", "本月", "季度", "截至", "至今")
_METRIC_MARKERS = ("收益率", "最大回撤", "回撤", "收益", "波动率", "夏普")
def build_clarification(question: str) -> dict | None:
"""识别缺少统计时间范围的指标问题,返回结构化反问。"""
text = (question or "").strip()
if any(marker in text for marker in _METRIC_MARKERS) and not any(
marker in text for marker in _TIME_MARKERS
):
return {
"required": True,
"missing": ["time_range"],
"message": "请补充收益率的统计时间范围。",
}
return None
def apply_query_options(
sql: str,
*,
page: int = 1,
page_size: int = 100,
sort_by: str | None = None,
sort_order: str = "asc",
) -> str:
"""通过 SQL AST 添加分页和白名单标识符排序。"""
if page < 1 or page_size < 1:
raise ValueError("分页参数必须为正数")
if sort_order not in {"asc", "desc"}:
raise ValueError("排序方向无效")
if sort_by is not None and not _IDENTIFIER.fullmatch(sort_by):
raise ValueError("排序字段无效")
statement = parse_one(sql, read="mysql")
if sort_by:
statement = statement.order_by(
exp.Ordered(this=exp.column(sort_by), desc=sort_order == "desc")
)
statement = statement.limit(page_size)
if page > 1:
statement = statement.offset((page - 1) * page_size)
return statement.sql(dialect="mysql")
def render_csv(columns: list[str], rows: list[dict]) -> str:
"""将已脱敏的查询结果渲染为 CSV。"""
output = StringIO()
writer = csv.DictWriter(output, fieldnames=columns)
writer.writeheader()
writer.writerows({column: row.get(column) for column in columns} for row in rows)
return output.getvalue()
def build_query_explanation(question: str, sql: str) -> dict:
"""返回指标口径和不含 SQL 文本的查询计划摘要。"""
statement = parse_one(sql, read="mysql")
tables = sorted({table.name for table in statement.find_all(exp.Table)})
return {
"metrics": resolve_semantics(question).get("metrics", []),
"plan": {
"tables": tables,
"operation": "SELECT",
"join_count": max(0, len(tables) - 1),
},
}