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
+66
View File
@@ -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