feat: add governed public knowledge retrieval

This commit is contained in:
张胜宇
2026-09-10 18:42:44 +08:00
parent ba22a2220f
commit 511a8ca18f
20 changed files with 729 additions and 7 deletions
+2
View File
@@ -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)
+43
View File
@@ -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, ...] = ()
+2 -1
View File
@@ -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
+28
View File
@@ -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)
+11
View File
@@ -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(
+100
View File
@@ -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
+11
View File
@@ -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)))
+55
View File
@@ -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)
+2 -1
View File
@@ -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)