feat:新增投顾agent和nl2sqlagent
This commit is contained in:
@@ -0,0 +1,66 @@
|
||||
"""NL2SQL Few-shot 示例召回。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from nl2sql.embedding import embed_texts
|
||||
from nl2sql.milvus_collections import NL2SQL_COLLECTION
|
||||
|
||||
|
||||
logger = logging.getLogger("nl2sql.few_shot")
|
||||
|
||||
|
||||
def _flatten(results: Any):
|
||||
for batch in results or []:
|
||||
if isinstance(batch, dict):
|
||||
yield batch
|
||||
else:
|
||||
yield from batch or []
|
||||
|
||||
|
||||
async def retrieve_few_shot(
|
||||
query: str,
|
||||
milvus_client,
|
||||
*,
|
||||
embedder: Callable[[list[str]], Awaitable[list[list[float]]]] = embed_texts,
|
||||
top_k: int = 3,
|
||||
threshold: float = 0.75,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""召回已标记为有效的 Few-shot 示例,异常时返回空列表。"""
|
||||
if not query or not query.strip() or top_k <= 0:
|
||||
return []
|
||||
try:
|
||||
vector = (await embedder([query]))[0]
|
||||
hits = await milvus_client.search(
|
||||
collection_name=NL2SQL_COLLECTION,
|
||||
data=[vector],
|
||||
limit=top_k,
|
||||
filter='chunk_type == "few_shot_example" and is_valid == true and is_deprecated == false',
|
||||
output_fields=["query", "correct_sql", "explanation", "case_id", "is_valid", "is_deprecated"],
|
||||
)
|
||||
except Exception: # noqa: BLE001 Few-shot 失败不阻断主查询
|
||||
logger.warning("NL2SQL Few-shot 召回失败", exc_info=True)
|
||||
return []
|
||||
|
||||
examples = []
|
||||
for hit in _flatten(hits):
|
||||
entity = hit.get("entity") or hit
|
||||
if not entity.get("is_valid", True) or entity.get("is_deprecated", False):
|
||||
continue
|
||||
distance = hit.get("distance")
|
||||
score = 1.0 - distance if distance is not None else hit.get("score")
|
||||
if score is None or score < threshold:
|
||||
continue
|
||||
example = {
|
||||
"query": entity.get("query", ""),
|
||||
"correct_sql": entity.get("correct_sql", ""),
|
||||
"score": score,
|
||||
}
|
||||
if entity.get("explanation"):
|
||||
example["explanation"] = entity["explanation"]
|
||||
if entity.get("case_id"):
|
||||
example["case_id"] = entity["case_id"]
|
||||
examples.append(example)
|
||||
return examples
|
||||
Reference in New Issue
Block a user