feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -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
|
||||
Reference in New Issue
Block a user