Files
Mutual_Fund/service/nl2sql/query_service.py
T

216 lines
7.9 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.query_rewriter import rewrite_query
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):
"""查询编排失败或当前用户没有查询权限。"""
def _filter_schema_columns(schema: dict[str, Any], permission: dict[str, Any]) -> dict[str, Any]:
"""只把当前用户有列权限的字段暴露给 SQL 生成模型。"""
authorized_columns = permission.get("columns")
if authorized_columns is None:
return schema
filtered = dict(schema)
filtered["columns"] = [
column
for column in schema.get("columns", [])
if column.get("field_name")
in set(authorized_columns.get(column.get("table_name"), set()))
]
return filtered
def _normalize_advisor_placeholder(sql: str) -> str:
"""统一 LLM 可能生成的投顾参数占位符,供 SQLAlchemy 命名绑定。"""
if "advisor_id" not in sql or "?" not in sql:
return sql
return sql.replace("?", ":advisor_id", 1)
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。"""
effective_question = await rewrite_query(
request.question,
conversation_context,
llm_client=llm_client,
)
try:
ensure_query_intent(effective_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(effective_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 = _filter_schema_columns(
await schema_loader(candidate_tables, permission),
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(effective_question)
except Exception: # noqa: BLE001 Few-shot 故障不阻断主查询
few_shot = []
generated = await generate_sql(
effective_question,
schema,
llm_client=llm_client,
few_shot=few_shot,
conversation_context=conversation_context,
)
generated = replace(generated, sql=_normalize_advisor_placeholder(generated.sql))
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