feat:新增投顾agent和nl2sqlagent

This commit is contained in:
2026-09-13 16:19:24 +08:00
parent c80c6acac0
commit 163192bf55
122 changed files with 7488 additions and 362 deletions
+92
View File
@@ -0,0 +1,92 @@
"""NL2SQL 元数据召回与授权过滤。"""
from __future__ import annotations
from collections.abc import Awaitable, Callable, Iterable
from typing import Any
from nl2sql.embedding import embed_texts
from nl2sql.milvus_collections import NL2SQL_COLLECTION
def _flatten_hits(results: Any) -> Iterable[dict[str, Any]]:
"""兼容 Milvus 返回的批次列表和单条字典结构。"""
for batch in results or []:
if isinstance(batch, dict):
yield batch
else:
yield from batch or []
def filter_authorized_metadata(
tables: list[dict[str, Any]],
columns: list[dict[str, Any]],
permission: dict[str, Any],
) -> dict[str, list[dict[str, Any]]]:
"""过滤失效实体、无表权限实体和无字段权限实体。"""
authorized_tables = set(permission.get("tables", set()))
authorized_columns = permission.get("columns", {})
valid_table_names = {
table.get("table_name")
for table in tables
if table.get("is_valid", True) and table.get("table_name") in authorized_tables
}
filtered_tables = [
table
for table in tables
if table.get("is_valid", True)
and table.get("table_name") in valid_table_names
]
filtered_columns = [
column
for column in columns
if column.get("is_valid", True)
and column.get("table_name") in valid_table_names
and column.get("field_name")
in set(authorized_columns.get(column.get("table_name"), set()))
]
return {"tables": filtered_tables, "columns": filtered_columns}
async def retrieve_metadata(
query: str,
milvus_client,
*,
embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts,
top_k: int = 5,
) -> list[dict[str, Any]]:
"""分别召回表级和字段级元数据,并过滤无效向量。"""
if not query or not query.strip():
return []
if top_k <= 0:
return []
vector = (await embedder([query]))[0]
output_fields = [
"id",
"chunk_type",
"table_name",
"field_name",
"text",
"is_valid",
"is_deprecated",
]
results: list[dict[str, Any]] = []
for chunk_type in ("table_meta", "field_meta"):
hits = await milvus_client.search(
collection_name=NL2SQL_COLLECTION,
data=[vector],
limit=top_k,
filter=f'chunk_type == "{chunk_type}" and is_valid == true',
output_fields=output_fields,
)
for hit in _flatten_hits(hits):
entity = hit.get("entity") or hit
if not entity.get("is_valid", True) or entity.get("is_deprecated", False):
continue
results.append(
{
**entity,
"distance": hit.get("distance", hit.get("score")),
}
)
return results