78 lines
2.6 KiB
Python
78 lines
2.6 KiB
Python
"""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),
|
||
|
|
},
|
||
|
|
}
|