chore: update gitignore; feat: 新增customer_agent业务模块与api路由
This commit is contained in:
+158
@@ -0,0 +1,158 @@
|
||||
"""Milvus-only retrieval primitives for the customer-service RAG layer."""
|
||||
from __future__ import annotations
|
||||
|
||||
from inspect import isawaitable
|
||||
import logging
|
||||
|
||||
from rag.embedding import embed_texts
|
||||
|
||||
|
||||
logger = logging.getLogger("rag.retrieve")
|
||||
|
||||
|
||||
BUSINESS_RETRIEVAL = (
|
||||
("fin_faq", "faq", 3, 0.75),
|
||||
("fin_fund_doc", "funddoc", 5, 0.70),
|
||||
("fin_policy", "policy", 5, 0.70),
|
||||
)
|
||||
|
||||
|
||||
async def _config_value(config_getter, key: str, default):
|
||||
value = config_getter(key, str(default))
|
||||
if isawaitable(value):
|
||||
value = await value
|
||||
return type(default)(value)
|
||||
|
||||
|
||||
async def retrieve_candidates(
|
||||
query: str,
|
||||
customer_id: str | None,
|
||||
*,
|
||||
milvus_client,
|
||||
embedder=embed_texts,
|
||||
config_getter,
|
||||
) -> list[dict]:
|
||||
if not query or not query.strip():
|
||||
return []
|
||||
vectors = await embedder([query])
|
||||
vector = vectors[0]
|
||||
plans = list(BUSINESS_RETRIEVAL)
|
||||
if customer_id:
|
||||
plans.append(("customer_memory", "memory", 5, 0.60))
|
||||
|
||||
candidates = []
|
||||
for collection_name, key_suffix, default_topk, default_threshold in plans:
|
||||
topk = await _config_value(
|
||||
config_getter,
|
||||
f"agent.customer.rag.topk.{key_suffix}",
|
||||
default_topk,
|
||||
)
|
||||
threshold = await _config_value(
|
||||
config_getter,
|
||||
f"agent.customer.rag.threshold.{key_suffix}",
|
||||
default_threshold,
|
||||
)
|
||||
expression = ""
|
||||
if collection_name == "customer_memory":
|
||||
escaped = customer_id.replace("\\", "\\\\").replace('"', '\\"')
|
||||
expression = f'customer_id == "{escaped}"'
|
||||
result = await milvus_client.search(
|
||||
collection_name=collection_name,
|
||||
data=[vector],
|
||||
limit=topk,
|
||||
filter=expression,
|
||||
output_fields=["doc_id", "title", "section_title", "text", "strategy"],
|
||||
)
|
||||
candidates.append(
|
||||
{
|
||||
"collection_name": collection_name,
|
||||
"threshold": threshold,
|
||||
"results": result,
|
||||
}
|
||||
)
|
||||
return candidates
|
||||
|
||||
|
||||
def _flatten_hits(results):
|
||||
for batch in results or []:
|
||||
if isinstance(batch, dict):
|
||||
yield batch
|
||||
else:
|
||||
yield from batch or []
|
||||
|
||||
|
||||
def _format_candidates(candidates) -> list[dict]:
|
||||
sources = []
|
||||
for candidate in candidates:
|
||||
threshold = candidate["threshold"]
|
||||
for hit in _flatten_hits(candidate["results"]):
|
||||
entity = hit.get("entity") or hit
|
||||
score = hit.get("distance", hit.get("score"))
|
||||
if score is None or score < threshold:
|
||||
continue
|
||||
sources.append(
|
||||
{
|
||||
"doc_id": entity.get("doc_id", ""),
|
||||
"title": entity.get("title", ""),
|
||||
"section_title": entity.get("section_title") or None,
|
||||
"chunk_text": entity.get("text", ""),
|
||||
"score": score,
|
||||
}
|
||||
)
|
||||
sources.sort(key=lambda source: source["score"], reverse=True)
|
||||
return sources
|
||||
|
||||
|
||||
async def rag_retrieve(
|
||||
query: str,
|
||||
customer_id: str | None,
|
||||
*,
|
||||
milvus_client,
|
||||
embedder=embed_texts,
|
||||
config_getter,
|
||||
) -> list[dict]:
|
||||
"""Return the stable source contract consumed by客服 Agent."""
|
||||
try:
|
||||
candidates = await retrieve_candidates(
|
||||
query,
|
||||
customer_id,
|
||||
milvus_client=milvus_client,
|
||||
embedder=embedder,
|
||||
config_getter=config_getter,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("RAG retrieval failed")
|
||||
return []
|
||||
return _format_candidates(candidates)
|
||||
|
||||
|
||||
async def retrieve_with_status(
|
||||
query: str,
|
||||
customer_id: str | None,
|
||||
*,
|
||||
milvus_client,
|
||||
embedder=embed_texts,
|
||||
config_getter,
|
||||
) -> dict:
|
||||
"""Expose operational status while keeping failed source lists empty."""
|
||||
try:
|
||||
vectors = await embedder([query])
|
||||
except Exception:
|
||||
logger.exception("RAG embedding failed")
|
||||
return {"status": "embedding_failed", "sources": []}
|
||||
|
||||
async def reuse_vector(_texts):
|
||||
return vectors
|
||||
|
||||
try:
|
||||
candidates = await retrieve_candidates(
|
||||
query,
|
||||
customer_id,
|
||||
milvus_client=milvus_client,
|
||||
embedder=reuse_vector,
|
||||
config_getter=config_getter,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Milvus retrieval failed")
|
||||
return {"status": "milvus_unavailable", "sources": []}
|
||||
return {"status": "ok", "sources": _format_candidates(candidates)}
|
||||
Reference in New Issue
Block a user