feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,77 @@
|
||||
"""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),
|
||||
},
|
||||
}
|
||||
Reference in New Issue
Block a user