feat: add governed public knowledge retrieval
This commit is contained in:
@@ -54,6 +54,8 @@ class Settings(BaseSettings):
|
||||
model_router_config_ref: str = "local"
|
||||
model_default_endpoint: str = ""
|
||||
model_fallback_endpoint: str = ""
|
||||
knowledge_embedding_endpoint_code: str = ""
|
||||
knowledge_embedding_timeout_ms: int = Field(default=15000, gt=0)
|
||||
sse_heartbeat_seconds: int = Field(default=15, gt=0)
|
||||
sse_chunk_characters: int = Field(default=256, ge=1, le=4096)
|
||||
sse_max_connection_seconds: int = Field(default=300, gt=0)
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
ALLOWED_KNOWLEDGE_COLLECTIONS = frozenset({
|
||||
"fin_faq_collection",
|
||||
"fin_product_collection",
|
||||
"fin_policy_collection",
|
||||
})
|
||||
|
||||
|
||||
class KnowledgeQuery(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
query: str = Field(min_length=1, max_length=2000)
|
||||
intents: tuple[str, ...] = Field(min_length=1, max_length=4)
|
||||
top_k: int = Field(default=5, ge=1, le=20)
|
||||
|
||||
@field_validator("query")
|
||||
@classmethod
|
||||
def query_must_not_be_blank(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("query must not be blank")
|
||||
return value
|
||||
|
||||
|
||||
class KnowledgeHit(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
knowledge_id: str
|
||||
collection: str
|
||||
snippet: str
|
||||
title: str | None = None
|
||||
answer: str | None = None
|
||||
score: float | None = Field(default=None, ge=0, le=1)
|
||||
version: str | None = None
|
||||
|
||||
|
||||
class KnowledgeSearchResult(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", frozen=True)
|
||||
|
||||
hits: tuple[KnowledgeHit, ...] = ()
|
||||
degraded: bool = False
|
||||
degradation_reason: str | None = None
|
||||
searched_collections: tuple[str, ...] = ()
|
||||
@@ -92,6 +92,7 @@ class JwtAuthenticator:
|
||||
if claims.get("visitor") is True:
|
||||
return RequestContext(
|
||||
user_id=str(subject), trace_id=str(uuid4()), roles=("visitor",),
|
||||
permissions=("agent:run",), data_scope="public",
|
||||
# 访客仅可运行 Agent 与读取已发布的公共知识,绝不含个人数据权限。
|
||||
permissions=("agent:run", "knowledge:query"), data_scope="public",
|
||||
)
|
||||
return RequestContext(user_id=str(claims["sub"]), trace_id=str(uuid4()))
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
from typing import Any
|
||||
|
||||
from app.core.errors import ForbiddenAgentError, RecoverableAgentError
|
||||
from app.core.knowledge_contracts import ALLOWED_KNOWLEDGE_COLLECTIONS
|
||||
|
||||
|
||||
class MilvusKnowledgeClient:
|
||||
def __init__(self, uri: str, token: str | None = None) -> None:
|
||||
self._uri = uri
|
||||
self._token = token
|
||||
self._client: Any | None = None
|
||||
|
||||
async def _ensure_client(self) -> Any:
|
||||
if self._client is None:
|
||||
from pymilvus import AsyncMilvusClient # type: ignore[import-untyped]
|
||||
|
||||
self._client = AsyncMilvusClient(uri=self._uri, token=self._token)
|
||||
return self._client
|
||||
|
||||
async def search(
|
||||
self, collection: str, vector: list[float], top_k: int
|
||||
) -> list[dict[str, Any]]:
|
||||
if collection not in ALLOWED_KNOWLEDGE_COLLECTIONS:
|
||||
raise ForbiddenAgentError("未授权的知识集合")
|
||||
if len(vector) != 1024 or not 1 <= top_k <= 20:
|
||||
raise RecoverableAgentError("知识检索参数无效")
|
||||
try:
|
||||
client = await self._ensure_client()
|
||||
batches = await client.search(
|
||||
collection_name=collection,
|
||||
data=[vector],
|
||||
limit=top_k,
|
||||
output_fields=["knowledge_id", "title", "snippet", "tags", "version"],
|
||||
search_params={"metric_type": "COSINE"},
|
||||
)
|
||||
except Exception as exc:
|
||||
raise RecoverableAgentError("知识检索不可用") from exc
|
||||
return [
|
||||
normalized
|
||||
for batch in batches
|
||||
for hit in batch
|
||||
if (normalized := self._normalize_hit(hit)) is not None
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _normalize_hit(hit: Any) -> dict[str, Any] | None:
|
||||
"""统一 Milvus SDK 的平铺与 entity 包装命中格式。"""
|
||||
raw = dict(hit)
|
||||
entity = raw.get("entity")
|
||||
fields = entity if isinstance(entity, dict) else raw
|
||||
knowledge_id = fields.get("knowledge_id")
|
||||
snippet = fields.get("snippet")
|
||||
score = raw.get("score", raw.get("distance", fields.get("score")))
|
||||
if (
|
||||
not isinstance(knowledge_id, str)
|
||||
or not isinstance(snippet, str)
|
||||
or not isinstance(score, (int, float))
|
||||
or isinstance(score, bool)
|
||||
):
|
||||
return None
|
||||
normalized: dict[str, Any] = {
|
||||
"knowledge_id": knowledge_id,
|
||||
"snippet": snippet,
|
||||
"score": float(score),
|
||||
}
|
||||
for field in ("title", "tags", "version"):
|
||||
value = fields.get(field)
|
||||
if value is not None:
|
||||
normalized[field] = value
|
||||
return normalized
|
||||
@@ -0,0 +1,28 @@
|
||||
from datetime import date, datetime
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import JSON, BigInteger, Date, DateTime, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.model.base import Base
|
||||
|
||||
|
||||
class FinKnowledgeMeta(Base):
|
||||
__tablename__ = "fin_knowledge_meta"
|
||||
|
||||
id: Mapped[int] = mapped_column(BigInteger, primary_key=True)
|
||||
knowledge_type: Mapped[str] = mapped_column(String(32))
|
||||
title: Mapped[str] = mapped_column(String(256))
|
||||
source_file: Mapped[str | None] = mapped_column(String(256))
|
||||
minio_path: Mapped[str | None] = mapped_column(String(512))
|
||||
milvus_collection: Mapped[str] = mapped_column(String(64))
|
||||
version: Mapped[str | None] = mapped_column(String(16))
|
||||
effective_date: Mapped[date | None] = mapped_column(Date)
|
||||
expire_date: Mapped[date | None] = mapped_column(Date)
|
||||
content_text: Mapped[str] = mapped_column(Text)
|
||||
tags: Mapped[list[Any] | None] = mapped_column(JSON)
|
||||
reviewer_id: Mapped[int | None] = mapped_column(BigInteger)
|
||||
review_status: Mapped[str] = mapped_column(String(16))
|
||||
status: Mapped[str] = mapped_column(String(16))
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
updated_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -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)
|
||||
@@ -110,7 +110,8 @@ class WorkerRuntime:
|
||||
if actor_type == "visitor":
|
||||
return identity.model_copy(update={
|
||||
"roles": ("visitor",),
|
||||
"permissions": ("agent:run",),
|
||||
# 与访客 JWT 对齐,只恢复公开 Agent 和公开知识的最小权限。
|
||||
"permissions": ("agent:run", "knowledge:query"),
|
||||
"data_scope": "public",
|
||||
})
|
||||
return await self.resolve_identity(identity)
|
||||
|
||||
Reference in New Issue
Block a user