"""执行 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))