"""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), }, }