add .gitattributes: python文件统一LF换行
This commit is contained in:
@@ -10,6 +10,7 @@ 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
|
||||
@@ -19,6 +20,29 @@ 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,
|
||||
*,
|
||||
@@ -30,8 +54,13 @@ async def query(
|
||||
conversation_context: str = "",
|
||||
):
|
||||
"""执行权限、召回、权威 Schema、生成和安全校验,返回可执行 SQL。"""
|
||||
effective_question = await rewrite_query(
|
||||
request.question,
|
||||
conversation_context,
|
||||
llm_client=llm_client,
|
||||
)
|
||||
try:
|
||||
ensure_query_intent(request.question)
|
||||
ensure_query_intent(effective_question)
|
||||
except UnsupportedIntent as exc:
|
||||
raise QueryServiceError(str(exc)) from exc
|
||||
permission = await permission_loader(request.user_id)
|
||||
@@ -40,7 +69,7 @@ async def query(
|
||||
if metadata_retriever is None or schema_loader is None:
|
||||
raise QueryServiceError("NL2SQL 查询依赖未配置")
|
||||
|
||||
hits = await metadata_retriever(request.question)
|
||||
hits = await metadata_retriever(effective_question)
|
||||
candidate_tables = {
|
||||
hit.get("table_name")
|
||||
for hit in hits
|
||||
@@ -48,23 +77,27 @@ async def query(
|
||||
}
|
||||
if not candidate_tables:
|
||||
raise QueryServiceError("未找到有权限的业务表")
|
||||
schema = await schema_loader(candidate_tables, permission)
|
||||
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(request.question)
|
||||
few_shot = await few_shot_retriever(effective_question)
|
||||
except Exception: # noqa: BLE001 Few-shot 故障不阻断主查询
|
||||
few_shot = []
|
||||
generated = await generate_sql(
|
||||
request.question,
|
||||
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(
|
||||
|
||||
Reference in New Issue
Block a user