"""供其他 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"]) validated_option_sql = validate_select_sql( option_sql, authorized_tables=permission.get("tables", set()), authorized_columns=final_columns, max_rows=max_rows, ) return validated_option_sql 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: parameters = ( {"advisor_id": request.user_id} if ":advisor_id" in validated.sql else None ) 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, parameters=parameters, 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