add .gitattributes: python文件统一LF换行

This commit is contained in:
2026-09-14 13:00:15 +08:00
parent 67d5cfc2b8
commit bc7761d79e
6 changed files with 64 additions and 8 deletions
+38 -5
View File
@@ -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(