148 lines
5.6 KiB
Python
148 lines
5.6 KiB
Python
"""执行 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))
|