Files
Mutual_Fund/scripts/check_nl2sql_e2e.py
T

148 lines
5.6 KiB
Python
Raw Normal View History

2026-09-13 16:19:24 +08:00
"""执行 NL2SQL 真实依赖和元数据链路验收。"""
from __future__ import annotations
import asyncio
import argparse
import sys
from pathlib import Path
from uuid import uuid4
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from config import database
from config.database.milvus import client as milvus_client
from config.settings import settings
from nl2sql.contracts import DataQueryRequest
from nl2sql.health import check_nl2sql_health
from nl2sql.milvus_collections import NL2SQL_COLLECTION
from nl2sql.retrieval import retrieve_metadata
from nl2sql.schema import load_authoritative_schema
from service.nl2sql.query_service import execute_query
from service.nl2sql.permission_service import load_query_permission
from tool.llm import llm
async def run_real_query(user_id: int, question: str, *, query_id: str | None = None) -> dict:
"""使用已有员工权限执行一条真实查询,只输出脱敏统计。"""
from config.database.mysql import get_session_factory
query_id = query_id or uuid4().hex
try:
async with get_session_factory()() as db:
permission = await load_query_permission(db, user_id)
if not permission.get("can_query"):
return {"query_status": "permission_denied", "user_id": user_id}
async def permission_loader(_user_id: int):
return permission
async def metadata_retriever(text_value: str):
return await retrieve_metadata(text_value, milvus_client())
async def schema_loader(table_names: set[str], _permission: dict):
return await load_authoritative_schema(
db,
database=settings.mysql.database,
candidate_tables=table_names,
)
result = await execute_query(
DataQueryRequest(
question=question,
user_id=user_id,
trace_id=f"nl2sql-e2e-{query_id}",
include_sql=False,
),
session=db,
query_id=query_id,
permission_loader=permission_loader,
metadata_retriever=metadata_retriever,
schema_loader=schema_loader,
llm_client=llm,
summary_llm=llm,
masks=permission.get("masks"),
)
return {
"query_status": "success",
"row_count": result.row_count,
"truncated": result.truncated,
"columns": result.columns,
"has_summary": bool(result.summary),
}
except Exception as exc: # noqa: BLE001 验收脚本输出结构化失败摘要
return {"query_status": "failed", "error_type": type(exc).__name__}
def summarize_real_queries(results: list[dict]) -> dict:
"""汇总真实查询结果,只保留状态和数量统计。"""
success = [item for item in results if item.get("query_status") == "success"]
return {
"total": len(results),
"success": len(success),
"failed": sum(1 for item in results if item.get("query_status") == "failed"),
"non_empty_success": sum(1 for item in success if item.get("row_count", 0) > 0),
"all_success": bool(results) and len(success) == len(results),
}
async def main(
user_id: int | None = None,
question: str | None = None,
questions: list[str] | None = None,
) -> None:
health = await check_nl2sql_health(
{
"mysql": database.mysql.check_health,
"redis": database.redis.check_health,
"milvus": database.milvus.check_health,
"llm": llm.check_health,
}
)
client = milvus_client()
description = await client.describe_collection(collection_name=NL2SQL_COLLECTION)
rows = await client.query(
collection_name=NL2SQL_COLLECTION,
filter="is_valid == true and is_deprecated == false",
output_fields=["id"],
)
vector_dim = next(
field["params"]["dim"]
for field in description["fields"]
if field["name"] == "vector"
)
embedding_dim = None
embedding_error = None
try:
embedding_dim = len(await llm.embed_one("NL2SQL 验收测试"))
except Exception as exc: # noqa: BLE001 记录类型,不输出密钥或请求内容
embedding_error = type(exc).__name__
output = {
"health": health,
"collection": NL2SQL_COLLECTION,
"configured_dimension": settings.llm.embed_dimensions,
"collection_dimension": vector_dim,
"valid_vector_count": len(rows or []),
"embedding_dimension": embedding_dim,
"embedding_error": embedding_error,
}
query_questions = questions or ([question] if question else [])
if user_id is not None and query_questions:
query_results = [
await run_real_query(user_id, item, query_id=f"nl2sql-e2e-{index}-{uuid4().hex[:8]}")
for index, item in enumerate(query_questions, start=1)
]
output["queries"] = query_results
output["query_summary"] = summarize_real_queries(query_results)
if len(query_results) == 1:
output["query"] = query_results[0]
print(output)
await database.dispose()
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="NL2SQL 真实依赖和查询验收")
parser.add_argument("--user-id", type=int, help="已配置 NL2SQL 权限的员工 ID")
parser.add_argument("--question", action="append", help="真实业务查询问题,可重复传入")
args = parser.parse_args()
asyncio.run(main(args.user_id, questions=args.question))