feat: add governed public knowledge retrieval
This commit is contained in:
@@ -7,6 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from app.core.config import get_settings
|
||||
from app.core.errors import RecoverableAgentError
|
||||
from app.core.fund_contracts import FundQuoteQuery
|
||||
from app.core.knowledge_contracts import KnowledgeQuery
|
||||
from app.core.nl2sql_contracts import FinancialNL2SQLInput
|
||||
from app.infrastructure.fund_quote_cache import FundQuoteCache
|
||||
from app.infrastructure.memory_cache import MemoryCacheAdapter
|
||||
@@ -18,6 +19,7 @@ from app.service.agent.offsite_fund_agent import OffsiteFundAgent
|
||||
from app.service.financial_nl2sql_service import query_financial_data_tool
|
||||
from app.service.fund_quote_service import query_fund_quote_tool
|
||||
from app.service.intent_classifier import IntentClassifier
|
||||
from app.service.knowledge_tool_service import query_knowledge_tool
|
||||
from app.service.memory_recall_service import MemoryRecallService
|
||||
from app.service.model_gateway import (
|
||||
DatabaseModelEndpointResolver,
|
||||
@@ -154,6 +156,15 @@ def get_agent_factory() -> AgentFactory:
|
||||
allowed_roles=("advisor", "operator", "admin", "super_admin"),
|
||||
timeout_seconds=10,
|
||||
))
|
||||
registry.register(ToolDefinition(
|
||||
name="query_knowledge",
|
||||
input_model=KnowledgeQuery,
|
||||
handler=cast(Any, query_knowledge_tool),
|
||||
required_permission="knowledge:query",
|
||||
allowed_roles=("visitor", "customer"),
|
||||
# 15 秒向量端点预算外预留检索和权威回查时间,防止正常降级被工具层提前中断。
|
||||
timeout_seconds=20,
|
||||
))
|
||||
model_service = get_model_service()
|
||||
endpoint_resolver = DatabaseModelEndpointResolver()
|
||||
factory = AgentFactory(
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
import json
|
||||
from datetime import UTC, datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.sql.elements import ColumnElement
|
||||
|
||||
from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery
|
||||
from app.model.knowledge import FinKnowledgeMeta
|
||||
|
||||
|
||||
class KnowledgeMysqlAuthority:
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
self._session = session
|
||||
|
||||
async def filter_published(self, hits: tuple[KnowledgeHit, ...]) -> list[KnowledgeHit]:
|
||||
ids = tuple(int(hit.knowledge_id) for hit in hits if hit.knowledge_id.isdecimal())
|
||||
if not ids:
|
||||
return []
|
||||
rows = await self._session.scalars(
|
||||
select(FinKnowledgeMeta).where(
|
||||
FinKnowledgeMeta.id.in_(ids), *self._published_filters()
|
||||
)
|
||||
)
|
||||
approved = {str(row.id): row for row in rows}
|
||||
result: list[KnowledgeHit] = []
|
||||
for hit in hits:
|
||||
row = approved.get(hit.knowledge_id)
|
||||
if row is None:
|
||||
continue
|
||||
answer = self.extract_answer(row.content_text).strip()
|
||||
if answer:
|
||||
result.append(hit.model_copy(update={
|
||||
"answer": answer, "version": row.version, "title": row.title,
|
||||
}))
|
||||
return result
|
||||
|
||||
async def search_keyword(
|
||||
self, query: KnowledgeQuery, collections: tuple[str, ...], top_k: int
|
||||
) -> list[KnowledgeHit]:
|
||||
"""向量服务不可用时,在原授权集合内执行受限的只读关键词检索。"""
|
||||
keyword = self._keyword(query.query)
|
||||
if not collections or not keyword:
|
||||
return []
|
||||
# 显式转义 LIKE 通配符,避免用户输入扩大关键词降级的匹配范围。
|
||||
escaped_keyword = keyword.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
rows = await self._session.scalars(
|
||||
select(FinKnowledgeMeta)
|
||||
.where(
|
||||
FinKnowledgeMeta.milvus_collection.in_(collections),
|
||||
*self._published_filters(),
|
||||
FinKnowledgeMeta.content_text.like(f"%{escaped_keyword}%", escape="\\"),
|
||||
)
|
||||
.order_by(FinKnowledgeMeta.id.desc())
|
||||
.limit(top_k)
|
||||
)
|
||||
result: list[KnowledgeHit] = []
|
||||
for row in rows:
|
||||
answer = self.extract_answer(row.content_text).strip()
|
||||
if not answer:
|
||||
continue
|
||||
result.append(
|
||||
KnowledgeHit(
|
||||
knowledge_id=str(row.id),
|
||||
collection=row.milvus_collection,
|
||||
title=row.title,
|
||||
snippet=answer[:300],
|
||||
answer=answer,
|
||||
version=row.version,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _keyword(query: str) -> str:
|
||||
"""压缩空白并限制关键词长度,避免降级查询承载无界输入。"""
|
||||
return "".join(query.split())[:64]
|
||||
|
||||
@staticmethod
|
||||
def _published_filters() -> tuple[ColumnElement[bool], ...]:
|
||||
today = datetime.now(UTC).date()
|
||||
return (
|
||||
FinKnowledgeMeta.review_status == "published",
|
||||
FinKnowledgeMeta.status == "active",
|
||||
(FinKnowledgeMeta.effective_date.is_(None))
|
||||
| (FinKnowledgeMeta.effective_date <= today),
|
||||
(FinKnowledgeMeta.expire_date.is_(None))
|
||||
| (FinKnowledgeMeta.expire_date > today),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def extract_answer(content_text: str) -> str:
|
||||
try:
|
||||
payload = json.loads(content_text)
|
||||
except json.JSONDecodeError:
|
||||
return content_text
|
||||
answer = payload.get("answer") if isinstance(payload, dict) else None
|
||||
if isinstance(answer, str):
|
||||
return answer
|
||||
return content_text
|
||||
@@ -0,0 +1,11 @@
|
||||
class KnowledgeRuntimeConfig:
|
||||
DEFAULT_ROUTES = {
|
||||
"faq": ("fin_faq_collection", 3),
|
||||
"product_inquiry": ("fin_product_collection", 5),
|
||||
"policy_explain": ("fin_policy_collection", 5),
|
||||
}
|
||||
|
||||
def __init__(self, *, vector_dim: int = 1024, similarity_threshold: float = 0.70) -> None:
|
||||
self.routes = dict(self.DEFAULT_ROUTES)
|
||||
self.vector_dim = vector_dim
|
||||
self.similarity_threshold = similarity_threshold
|
||||
@@ -0,0 +1,66 @@
|
||||
from typing import Any, Protocol
|
||||
|
||||
from app.core.contracts import RequestContext
|
||||
from app.core.errors import RecoverableAgentError
|
||||
from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery, KnowledgeSearchResult
|
||||
from app.service.knowledge_config import KnowledgeRuntimeConfig
|
||||
|
||||
# ruff: noqa: E501
|
||||
|
||||
|
||||
class KnowledgeEmbedder(Protocol):
|
||||
async def embed(self, text: str) -> list[float]: ...
|
||||
|
||||
|
||||
class KnowledgeVectorStore(Protocol):
|
||||
async def search(self, collection: str, vector: list[float], top_k: int) -> list[dict[str, Any]]: ...
|
||||
|
||||
|
||||
class KnowledgeAuthority(Protocol):
|
||||
async def filter_published(self, hits: tuple[KnowledgeHit, ...]) -> list[KnowledgeHit]: ...
|
||||
|
||||
async def search_keyword(
|
||||
self, query: KnowledgeQuery, collections: tuple[str, ...], top_k: int
|
||||
) -> list[KnowledgeHit]: ...
|
||||
|
||||
|
||||
class KnowledgeRetrievalService:
|
||||
def __init__(
|
||||
self, embedder: KnowledgeEmbedder, vector_store: KnowledgeVectorStore,
|
||||
config: KnowledgeRuntimeConfig, authority: KnowledgeAuthority,
|
||||
) -> None:
|
||||
self._embedder = embedder
|
||||
self._vector_store = vector_store
|
||||
self._config = config
|
||||
self._authority = authority
|
||||
|
||||
async def search(self, query: KnowledgeQuery, context: RequestContext) -> KnowledgeSearchResult:
|
||||
del context
|
||||
targets = tuple(self._config.routes[intent] for intent in query.intents if intent in self._config.routes)
|
||||
if not targets:
|
||||
return KnowledgeSearchResult()
|
||||
collections: list[str] = []
|
||||
candidates: list[KnowledgeHit] = []
|
||||
try:
|
||||
vector = await self._embedder.embed(query.query)
|
||||
if len(vector) != self._config.vector_dim:
|
||||
raise RecoverableAgentError("嵌入维度与集合定义不一致")
|
||||
for collection, configured_top_k in targets:
|
||||
collections.append(collection)
|
||||
for raw in await self._vector_store.search(
|
||||
collection, vector, min(query.top_k, configured_top_k)
|
||||
):
|
||||
knowledge_id, snippet, score = raw.get("knowledge_id"), raw.get("snippet"), raw.get("score")
|
||||
if isinstance(knowledge_id, str) and isinstance(snippet, str) and isinstance(score, (int, float)):
|
||||
if not isinstance(score, bool) and self._config.similarity_threshold <= score <= 1:
|
||||
candidates.append(KnowledgeHit(knowledge_id=knowledge_id, collection=collection, snippet=snippet, score=float(score)))
|
||||
except RecoverableAgentError:
|
||||
fallback_collections = tuple(dict.fromkeys(collection for collection, _ in targets))
|
||||
fallback_top_k = max(min(query.top_k, configured_top_k) for _, configured_top_k in targets)
|
||||
fallback_hits = await self._authority.search_keyword(query, fallback_collections, fallback_top_k)
|
||||
return KnowledgeSearchResult(
|
||||
hits=tuple(fallback_hits), degraded=True, degradation_reason="milvus_unavailable",
|
||||
searched_collections=fallback_collections,
|
||||
)
|
||||
hits = await self._authority.filter_published(tuple(candidates))
|
||||
return KnowledgeSearchResult(hits=tuple(hits), searched_collections=tuple(dict.fromkeys(collections)))
|
||||
@@ -0,0 +1,55 @@
|
||||
from typing import Protocol
|
||||
|
||||
from app.core.config import get_settings
|
||||
from app.core.contracts import RequestContext
|
||||
from app.core.knowledge_contracts import KnowledgeQuery, KnowledgeSearchResult
|
||||
from app.infrastructure.db import SessionFactory
|
||||
from app.infrastructure.milvus_knowledge_adapter import MilvusKnowledgeClient
|
||||
from app.service.knowledge_authority import KnowledgeMysqlAuthority
|
||||
from app.service.knowledge_config import KnowledgeRuntimeConfig
|
||||
from app.service.knowledge_retrieval_service import KnowledgeRetrievalService
|
||||
from app.service.model_gateway import DatabaseModelGateway
|
||||
|
||||
|
||||
class EmbeddingGateway(Protocol):
|
||||
async def embed(self, *, endpoint_code: str, text: str, timeout_ms: int) -> list[float]: ...
|
||||
|
||||
|
||||
class DatabaseEmbeddingAdapter:
|
||||
def __init__(
|
||||
self, endpoint_code: str, timeout_ms: int, *, gateway: EmbeddingGateway
|
||||
) -> None:
|
||||
self._endpoint_code = endpoint_code
|
||||
self._timeout_ms = timeout_ms
|
||||
self._gateway = gateway
|
||||
|
||||
async def embed(self, text: str) -> list[float]:
|
||||
return await self._gateway.embed(
|
||||
endpoint_code=self._endpoint_code, text=text, timeout_ms=self._timeout_ms
|
||||
)
|
||||
|
||||
|
||||
async def query_knowledge_tool(
|
||||
arguments: KnowledgeQuery, context: RequestContext
|
||||
) -> KnowledgeSearchResult:
|
||||
settings = get_settings()
|
||||
if not settings.knowledge_embedding_endpoint_code:
|
||||
return KnowledgeSearchResult(
|
||||
degraded=True, degradation_reason="embedding_endpoint_unconfigured"
|
||||
)
|
||||
# 知识向量端点与默认聊天端点隔离,避免回答模型被误用于检索。
|
||||
embedder = DatabaseEmbeddingAdapter(
|
||||
settings.knowledge_embedding_endpoint_code,
|
||||
settings.knowledge_embedding_timeout_ms,
|
||||
gateway=DatabaseModelGateway(),
|
||||
)
|
||||
vector_store = MilvusKnowledgeClient(
|
||||
settings.milvus_uri, token=settings.milvus_token or None
|
||||
)
|
||||
# 权威元数据只读回查,确保对客答案始终来自已发布、有效的知识条目。
|
||||
async with SessionFactory() as session:
|
||||
authority = KnowledgeMysqlAuthority(session)
|
||||
service = KnowledgeRetrievalService(
|
||||
embedder, vector_store, KnowledgeRuntimeConfig(), authority
|
||||
)
|
||||
return await service.search(arguments, context)
|
||||
Reference in New Issue
Block a user