feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,147 @@
|
||||
"""执行 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))
|
||||
Reference in New Issue
Block a user