feat:修复投顾agent功能
This commit is contained in:
+10
-3
@@ -1,6 +1,7 @@
|
||||
"""NL2SQL 元数据召回与授权过滤。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
@@ -70,15 +71,21 @@ async def retrieve_metadata(
|
||||
"is_valid",
|
||||
"is_deprecated",
|
||||
]
|
||||
results: list[dict[str, Any]] = []
|
||||
for chunk_type in ("table_meta", "field_meta"):
|
||||
hits = await milvus_client.search(
|
||||
async def search_chunk(chunk_type: str):
|
||||
return 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,
|
||||
)
|
||||
|
||||
table_hits, field_hits = await asyncio.gather(
|
||||
search_chunk("table_meta"),
|
||||
search_chunk("field_meta"),
|
||||
)
|
||||
results: list[dict[str, Any]] = []
|
||||
for hits in (table_hits, field_hits):
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user