176 lines
6.5 KiB
Python
176 lines
6.5 KiB
Python
"""供其他 Agent 复用的 NL2SQL 查询编排入口。"""
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Awaitable, Callable
|
|
from dataclasses import replace
|
|
from typing import Any
|
|
|
|
from nl2sql.contracts import DataQueryRequest
|
|
from nl2sql.executor import QueryExecutionError, execute_readonly_sql
|
|
from nl2sql.row_scope import RowScopeError, apply_row_scope
|
|
from nl2sql.result import build_chart_config, summarize_result
|
|
from nl2sql.query_experience import apply_query_options, build_query_explanation
|
|
from nl2sql.supervisor import UnsupportedIntent, ensure_query_intent
|
|
from nl2sql.sql_agent import generate_sql
|
|
from nl2sql.sql_security import SqlSecurityError, validate_select_sql
|
|
|
|
|
|
class QueryServiceError(RuntimeError):
|
|
"""查询编排失败或当前用户没有查询权限。"""
|
|
|
|
|
|
async def query(
|
|
request: DataQueryRequest,
|
|
*,
|
|
permission_loader: Callable[[int], Awaitable[dict[str, Any]]],
|
|
metadata_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None,
|
|
schema_loader: Callable[[set[str], dict[str, Any]], Awaitable[dict[str, Any]]] | None,
|
|
llm_client=None,
|
|
few_shot_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None = None,
|
|
conversation_context: str = "",
|
|
):
|
|
"""执行权限、召回、权威 Schema、生成和安全校验,返回可执行 SQL。"""
|
|
try:
|
|
ensure_query_intent(request.question)
|
|
except UnsupportedIntent as exc:
|
|
raise QueryServiceError(str(exc)) from exc
|
|
permission = await permission_loader(request.user_id)
|
|
if not permission.get("can_query", False):
|
|
raise QueryServiceError("当前用户没有 NL2SQL 查询权限")
|
|
if metadata_retriever is None or schema_loader is None:
|
|
raise QueryServiceError("NL2SQL 查询依赖未配置")
|
|
|
|
hits = await metadata_retriever(request.question)
|
|
candidate_tables = {
|
|
hit.get("table_name")
|
|
for hit in hits
|
|
if hit.get("table_name") in permission.get("tables", set())
|
|
}
|
|
if not candidate_tables:
|
|
raise QueryServiceError("未找到有权限的业务表")
|
|
schema = await schema_loader(candidate_tables, permission)
|
|
if not schema.get("tables"):
|
|
raise QueryServiceError("候选表未通过权威 Schema 校验")
|
|
|
|
few_shot = []
|
|
if few_shot_retriever is not None:
|
|
try:
|
|
few_shot = await few_shot_retriever(request.question)
|
|
except Exception: # noqa: BLE001 Few-shot 故障不阻断主查询
|
|
few_shot = []
|
|
generated = await generate_sql(
|
|
request.question,
|
|
schema,
|
|
llm_client=llm_client,
|
|
few_shot=few_shot,
|
|
conversation_context=conversation_context,
|
|
)
|
|
max_rows = request.max_rows or permission.get("max_rows") or 1000
|
|
try:
|
|
validated = validate_select_sql(
|
|
generated.sql,
|
|
authorized_tables=permission.get("tables", set()),
|
|
authorized_columns=permission.get("columns"),
|
|
max_rows=max_rows,
|
|
)
|
|
try:
|
|
scoped_sql = apply_row_scope(
|
|
validated.sql,
|
|
permission,
|
|
request.data_scope,
|
|
)
|
|
except RowScopeError as exc:
|
|
raise QueryServiceError("行级权限范围无效") from exc
|
|
option_sql = apply_query_options(
|
|
scoped_sql,
|
|
page=request.page,
|
|
page_size=request.page_size,
|
|
sort_by=request.sort_by,
|
|
sort_order=request.sort_order,
|
|
)
|
|
# columns 为 None 表示不做列级限制;仅当配置了列权限时才需要
|
|
# 保证行级范围列可访问。保持 dict(含空 dict)行为不变。
|
|
if permission.get("columns") is None:
|
|
final_columns = None
|
|
else:
|
|
final_columns = {
|
|
table: set(columns)
|
|
for table, columns in permission["columns"].items()
|
|
}
|
|
for table, scope in (permission.get("row_scopes") or {}).items():
|
|
if scope.get("column"):
|
|
final_columns.setdefault(table, set()).add(scope["column"])
|
|
return validate_select_sql(
|
|
option_sql,
|
|
authorized_tables=permission.get("tables", set()),
|
|
authorized_columns=final_columns,
|
|
max_rows=max_rows,
|
|
)
|
|
except SqlSecurityError as exc:
|
|
raise QueryServiceError("生成的 SQL 未通过安全校验") from exc
|
|
|
|
|
|
async def execute_query(
|
|
request: DataQueryRequest,
|
|
*,
|
|
session,
|
|
query_id: str,
|
|
permission_loader: Callable[[int], Awaitable[dict[str, Any]]],
|
|
metadata_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None,
|
|
schema_loader: Callable[[set[str], dict[str, Any]], Awaitable[dict[str, Any]]] | None,
|
|
llm_client=None,
|
|
masks: dict[tuple[str, str], str] | None = None,
|
|
timeout_seconds: float | None = None,
|
|
summary_llm=None,
|
|
few_shot_retriever: Callable[[str], Awaitable[list[dict[str, Any]]]] | None = None,
|
|
conversation_context: str = "",
|
|
):
|
|
"""完成 SQL 编排、只读执行和统一结果返回。"""
|
|
validated = await query(
|
|
request,
|
|
permission_loader=permission_loader,
|
|
metadata_retriever=metadata_retriever,
|
|
schema_loader=schema_loader,
|
|
llm_client=llm_client,
|
|
few_shot_retriever=few_shot_retriever,
|
|
conversation_context=conversation_context,
|
|
)
|
|
try:
|
|
result = await execute_readonly_sql(
|
|
session,
|
|
validated,
|
|
query_id=query_id,
|
|
trace_id=request.trace_id,
|
|
user_id=request.user_id,
|
|
masks=masks,
|
|
max_rows=request.max_rows,
|
|
timeout_seconds=timeout_seconds,
|
|
)
|
|
except QueryExecutionError as exc:
|
|
raise QueryServiceError("查询执行失败") from exc
|
|
result = replace(
|
|
result,
|
|
summary=(
|
|
await summarize_result(
|
|
request.question,
|
|
result.columns,
|
|
result.rows,
|
|
llm_client=summary_llm,
|
|
)
|
|
if summary_llm is not None
|
|
else None
|
|
),
|
|
chart=build_chart_config(result.columns, result.rows),
|
|
metric_definitions=build_query_explanation(request.question, validated.sql)["metrics"],
|
|
query_plan={
|
|
**build_query_explanation(request.question, validated.sql)["plan"],
|
|
"page": request.page,
|
|
"page_size": request.page_size,
|
|
"sort_by": request.sort_by,
|
|
"sort_order": request.sort_order,
|
|
},
|
|
)
|
|
if not request.include_sql:
|
|
result = replace(result, sql=None)
|
|
return result
|