feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,170 @@
|
||||
"""供其他 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,
|
||||
)
|
||||
final_columns = {
|
||||
table: set(columns)
|
||||
for table, columns in (permission.get("columns") or {}).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
|
||||
Reference in New Issue
Block a user