Files
Mutual_Fund/service/nl2sql/query_service.py
T

183 lines
6.8 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""供其他 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,
)
2026-09-13 20:48:42 +08:00
# 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"])
2026-09-14 10:57:48 +08:00
validated_option_sql = validate_select_sql(
2026-09-13 16:19:24 +08:00
option_sql,
authorized_tables=permission.get("tables", set()),
authorized_columns=final_columns,
max_rows=max_rows,
)
2026-09-14 10:57:48 +08:00
return validated_option_sql
2026-09-13 16:19:24 +08:00
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:
2026-09-14 10:57:48 +08:00
parameters = (
{"advisor_id": request.user_id}
if ":advisor_id" in validated.sql
else None
)
2026-09-13 16:19:24 +08:00
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,
2026-09-14 10:57:48 +08:00
parameters=parameters,
2026-09-13 16:19:24 +08:00
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