feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""NL2SQL 查询服务。"""
|
||||
@@ -0,0 +1,105 @@
|
||||
"""NL2SQL 管理能力服务。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from model.nl2sql_permission import (
|
||||
Nl2SqlQueryRole,
|
||||
Nl2SqlRoleColumnPermission,
|
||||
Nl2SqlRoleTablePermission,
|
||||
Nl2SqlSensitiveField,
|
||||
)
|
||||
from repositories.nl2sql_permission import Nl2SqlPermissionRepo
|
||||
|
||||
VALID_MASK_TYPES = {"partial", "hash"}
|
||||
VALID_ROW_SCOPE_TYPES = {"none", "customer_ids", "product_ids"}
|
||||
|
||||
|
||||
def _clean_text(value: str, name: str) -> str:
|
||||
cleaned = value.strip()
|
||||
if not cleaned:
|
||||
raise ValueError(f"{name}不能为空")
|
||||
return cleaned
|
||||
|
||||
|
||||
def validate_mask_type(mask_type: str | None) -> str | None:
|
||||
if mask_type is not None and mask_type not in VALID_MASK_TYPES:
|
||||
raise ValueError("脱敏类型必须为 partial 或 hash")
|
||||
return mask_type
|
||||
|
||||
|
||||
def validate_table_permission(payload: dict) -> dict:
|
||||
permission = payload.get("permission", "SELECT")
|
||||
if permission != "SELECT":
|
||||
raise ValueError("表权限只允许 SELECT")
|
||||
row_scope_type = payload.get("row_scope_type", "none")
|
||||
if row_scope_type not in VALID_ROW_SCOPE_TYPES:
|
||||
raise ValueError("行级范围类型无效")
|
||||
return {
|
||||
"table_name": _clean_text(payload["table_name"], "表名"),
|
||||
"permission": "SELECT",
|
||||
"row_scope_type": row_scope_type,
|
||||
"row_scope_column": payload.get("row_scope_column"),
|
||||
"status": payload.get("status", "active"),
|
||||
}
|
||||
|
||||
|
||||
async def create_role(db, payload: dict, *, repo_factory=Nl2SqlPermissionRepo):
|
||||
values = {
|
||||
"role_code": _clean_text(payload["role_code"], "角色编码"),
|
||||
"role_name": _clean_text(payload["role_name"], "角色名称"),
|
||||
"employee_role": _clean_text(payload["employee_role"], "员工角色"),
|
||||
"can_query": bool(payload.get("can_query", False)),
|
||||
"max_rows": int(payload.get("max_rows", 1000)),
|
||||
"daily_quota": int(payload.get("daily_quota", 0)),
|
||||
"status": "active",
|
||||
}
|
||||
if values["max_rows"] <= 0 or values["daily_quota"] < 0:
|
||||
raise ValueError("配额参数无效")
|
||||
return await repo_factory(db).add_role(**values)
|
||||
|
||||
|
||||
def role_payload(role: Nl2SqlQueryRole) -> dict:
|
||||
return {
|
||||
"id": role.id,
|
||||
"role_code": role.role_code,
|
||||
"role_name": role.role_name,
|
||||
"employee_role": role.employee_role,
|
||||
"can_query": role.can_query,
|
||||
"max_rows": role.max_rows,
|
||||
"daily_quota": role.daily_quota,
|
||||
"status": role.status,
|
||||
}
|
||||
|
||||
|
||||
def table_permission_payload(item: Nl2SqlRoleTablePermission) -> dict:
|
||||
return {
|
||||
"id": item.id,
|
||||
"role_id": item.role_id,
|
||||
"table_name": item.table_name,
|
||||
"permission": item.permission,
|
||||
"row_scope_type": item.row_scope_type,
|
||||
"row_scope_column": item.row_scope_column,
|
||||
"status": item.status,
|
||||
}
|
||||
|
||||
|
||||
def column_permission_payload(item: Nl2SqlRoleColumnPermission) -> dict:
|
||||
return {
|
||||
"id": item.id,
|
||||
"role_id": item.role_id,
|
||||
"table_name": item.table_name,
|
||||
"column_name": item.column_name,
|
||||
"access_mode": item.access_mode,
|
||||
"mask_type": item.mask_type,
|
||||
"status": item.status,
|
||||
}
|
||||
|
||||
|
||||
def sensitive_field_payload(item: Nl2SqlSensitiveField) -> dict:
|
||||
return {
|
||||
"id": item.id,
|
||||
"table_name": item.table_name,
|
||||
"column_name": item.column_name,
|
||||
"mask_type": item.mask_type,
|
||||
"description": item.description,
|
||||
"status": item.status,
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
"""NL2SQL 请求级权限快照服务。"""
|
||||
from __future__ import annotations
|
||||
|
||||
from nl2sql.permission import build_query_permission
|
||||
from model.sys_user import SysUser
|
||||
from repositories.nl2sql_permission import Nl2SqlPermissionRepo
|
||||
from repositories.sys_user import SysUserRepo
|
||||
|
||||
|
||||
def _denied_permission(user_id: int) -> dict:
|
||||
"""构造默认拒绝快照,避免把不存在用户的信息暴露给调用方。"""
|
||||
return build_query_permission(
|
||||
SysUser(id=user_id, user_type="UNKNOWN", employee_role=None, status="异常"),
|
||||
None,
|
||||
[],
|
||||
[],
|
||||
[],
|
||||
)
|
||||
|
||||
|
||||
async def load_query_permission(
|
||||
db,
|
||||
user_id: int,
|
||||
*,
|
||||
user_repo_factory=SysUserRepo,
|
||||
permission_repo_factory=Nl2SqlPermissionRepo,
|
||||
) -> dict:
|
||||
"""每次调用重新加载用户和 NL2SQL 权限,返回请求级权限快照。"""
|
||||
user = await user_repo_factory(db).get(user_id)
|
||||
if user is None or getattr(user, "status", None) != "正常":
|
||||
return _denied_permission(user_id)
|
||||
if getattr(user, "user_type", None) != "EMPLOYEE":
|
||||
return _denied_permission(user_id)
|
||||
|
||||
repo = permission_repo_factory(db)
|
||||
role = await repo.get_role_by_employee_role(user.employee_role)
|
||||
table_permissions = await repo.list_table_permissions(role.id) if role else []
|
||||
column_permissions = await repo.list_column_permissions(role.id) if role else []
|
||||
sensitive_fields = await repo.list_sensitive_fields()
|
||||
return build_query_permission(
|
||||
user,
|
||||
role,
|
||||
table_permissions,
|
||||
column_permissions,
|
||||
sensitive_fields,
|
||||
)
|
||||
@@ -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