From 511a8ca18f7e10644ef06efb2d006a9b7ac54d12 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=BC=A0=E8=83=9C=E5=AE=87?= <17412268+zzzzz11122222@user.noreply.gitee.com> Date: Thu, 10 Sep 2026 18:42:44 +0800 Subject: [PATCH] feat: add governed public knowledge retrieval --- .env.example | 2 + app/core/config.py | 2 + app/core/knowledge_contracts.py | 43 ++++++ app/core/security.py | 3 +- .../milvus_knowledge_adapter.py | 70 ++++++++++ app/model/knowledge.py | 28 ++++ app/service/agent/bootstrap.py | 11 ++ app/service/knowledge_authority.py | 100 +++++++++++++ app/service/knowledge_config.py | 11 ++ app/service/knowledge_retrieval_service.py | 66 +++++++++ app/service/knowledge_tool_service.py | 55 ++++++++ app/worker/runtime.py | 3 +- tests/unit/core/test_security.py | 4 +- .../test_milvus_knowledge_adapter.py | 47 +++++++ .../repository/test_fund_readonly_contract.py | 6 +- tests/unit/service/test_bootstrap.py | 4 + .../unit/service/test_knowledge_authority.py | 61 ++++++++ .../unit/service/test_knowledge_retrieval.py | 87 ++++++++++++ .../service/test_knowledge_tool_service.py | 131 ++++++++++++++++++ .../worker/test_runtime_worker_dispatch.py | 2 +- 20 files changed, 729 insertions(+), 7 deletions(-) create mode 100644 app/core/knowledge_contracts.py create mode 100644 app/infrastructure/milvus_knowledge_adapter.py create mode 100644 app/model/knowledge.py create mode 100644 app/service/knowledge_authority.py create mode 100644 app/service/knowledge_config.py create mode 100644 app/service/knowledge_retrieval_service.py create mode 100644 app/service/knowledge_tool_service.py create mode 100644 tests/unit/infrastructure/test_milvus_knowledge_adapter.py create mode 100644 tests/unit/service/test_knowledge_authority.py create mode 100644 tests/unit/service/test_knowledge_retrieval.py create mode 100644 tests/unit/service/test_knowledge_tool_service.py diff --git a/.env.example b/.env.example index f5917a6..e13338c 100644 --- a/.env.example +++ b/.env.example @@ -36,6 +36,8 @@ NEO4J_PASSWORD= MODEL_ROUTER_CONFIG_REF=local MODEL_DEFAULT_ENDPOINT= MODEL_FALLBACK_ENDPOINT= +KNOWLEDGE_EMBEDDING_ENDPOINT_CODE= +KNOWLEDGE_EMBEDDING_TIMEOUT_MS=15000 SSE_HEARTBEAT_SECONDS=15 SSE_MAX_CONNECTION_SECONDS=300 diff --git a/app/core/config.py b/app/core/config.py index 66250a3..d745d90 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -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) diff --git a/app/core/knowledge_contracts.py b/app/core/knowledge_contracts.py new file mode 100644 index 0000000..052aa05 --- /dev/null +++ b/app/core/knowledge_contracts.py @@ -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, ...] = () diff --git a/app/core/security.py b/app/core/security.py index 2fa9871..fa1f5e6 100644 --- a/app/core/security.py +++ b/app/core/security.py @@ -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())) diff --git a/app/infrastructure/milvus_knowledge_adapter.py b/app/infrastructure/milvus_knowledge_adapter.py new file mode 100644 index 0000000..98a9530 --- /dev/null +++ b/app/infrastructure/milvus_knowledge_adapter.py @@ -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 diff --git a/app/model/knowledge.py b/app/model/knowledge.py new file mode 100644 index 0000000..7306bdd --- /dev/null +++ b/app/model/knowledge.py @@ -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) diff --git a/app/service/agent/bootstrap.py b/app/service/agent/bootstrap.py index affabc7..b2ad3bc 100644 --- a/app/service/agent/bootstrap.py +++ b/app/service/agent/bootstrap.py @@ -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( diff --git a/app/service/knowledge_authority.py b/app/service/knowledge_authority.py new file mode 100644 index 0000000..73360ad --- /dev/null +++ b/app/service/knowledge_authority.py @@ -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 diff --git a/app/service/knowledge_config.py b/app/service/knowledge_config.py new file mode 100644 index 0000000..ca906ba --- /dev/null +++ b/app/service/knowledge_config.py @@ -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 diff --git a/app/service/knowledge_retrieval_service.py b/app/service/knowledge_retrieval_service.py new file mode 100644 index 0000000..54327a8 --- /dev/null +++ b/app/service/knowledge_retrieval_service.py @@ -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))) diff --git a/app/service/knowledge_tool_service.py b/app/service/knowledge_tool_service.py new file mode 100644 index 0000000..20eaa65 --- /dev/null +++ b/app/service/knowledge_tool_service.py @@ -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) diff --git a/app/worker/runtime.py b/app/worker/runtime.py index 274d9d7..98ed6c9 100644 --- a/app/worker/runtime.py +++ b/app/worker/runtime.py @@ -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) diff --git a/tests/unit/core/test_security.py b/tests/unit/core/test_security.py index d5d6722..722a211 100644 --- a/tests/unit/core/test_security.py +++ b/tests/unit/core/test_security.py @@ -53,7 +53,7 @@ def test_authenticate_visitor_token_returns_limited_anonymous_context() -> None: context = JwtAuthenticator(_settings()).authenticate(token) assert context.roles == ("visitor",) - assert context.permissions == ("agent:run",) + assert context.permissions == ("agent:run", "knowledge:query") assert context.customer_ids == () assert context.data_scope == "public" @@ -64,7 +64,7 @@ def test_visitor_token_issuer_creates_short_lived_limited_token() -> None: context = JwtAuthenticator(_settings()).authenticate(token) assert context.roles == ("visitor",) - assert context.permissions == ("agent:run",) + assert context.permissions == ("agent:run", "knowledge:query") assert expires_at > datetime.now(UTC) diff --git a/tests/unit/infrastructure/test_milvus_knowledge_adapter.py b/tests/unit/infrastructure/test_milvus_knowledge_adapter.py new file mode 100644 index 0000000..f730288 --- /dev/null +++ b/tests/unit/infrastructure/test_milvus_knowledge_adapter.py @@ -0,0 +1,47 @@ +import pytest + +from app.core.errors import ForbiddenAgentError +from app.infrastructure.milvus_knowledge_adapter import MilvusKnowledgeClient + + +class FakeMilvus: + def __init__(self) -> None: + self.kwargs = None + + async def search(self, **kwargs): + self.kwargs = kwargs + return [[{ + "distance": 0.91, + "entity": { + "knowledge_id": "101", + "snippet": "开户说明", + "title": "基金开户", + "tags": ["开户"], + "version": "v1", + }, + }]] + + +@pytest.mark.asyncio +async def test_knowledge_adapter_uses_cosine_and_minimal_public_projection() -> None: + client = MilvusKnowledgeClient("http://unused") + fake = FakeMilvus() + client._client = fake + + hits = await client.search("fin_faq_collection", [0.1] * 1024, 3) + + assert hits[0]["knowledge_id"] == "101" + assert hits[0]["snippet"] == "开户说明" + assert hits[0]["score"] == 0.91 + assert fake.kwargs["collection_name"] == "fin_faq_collection" + assert fake.kwargs["limit"] == 3 + assert fake.kwargs["search_params"] == {"metric_type": "COSINE"} + assert fake.kwargs["output_fields"] == ["knowledge_id", "title", "snippet", "tags", "version"] + + +@pytest.mark.asyncio +async def test_knowledge_adapter_rejects_non_public_collection() -> None: + client = MilvusKnowledgeClient("http://unused") + + with pytest.raises(ForbiddenAgentError): + await client.search("customer_vectors", [0.1] * 1024, 3) diff --git a/tests/unit/repository/test_fund_readonly_contract.py b/tests/unit/repository/test_fund_readonly_contract.py index 6f014cb..91bb2ba 100644 --- a/tests/unit/repository/test_fund_readonly_contract.py +++ b/tests/unit/repository/test_fund_readonly_contract.py @@ -9,6 +9,7 @@ from typing import Any import pytest +import app.model.fund as fund_models from app.model.base import Base REPOSITORY_PATH = "app/repository/fund_query_repository.py" @@ -123,10 +124,11 @@ def test_fund_model_module_contains_only_column_mappings() -> None: def test_fund_models_cover_all_fin_tables_with_expected_columns() -> None: + assert fund_models tables: dict[str, Any] = { name: table for name, table in Base.metadata.tables.items() - if name.startswith("fin_") + if name in FUND_TABLES } assert set(tables) == set(FUND_TABLES) for name, (column_count, primary_key) in FUND_TABLES.items(): @@ -137,7 +139,7 @@ def test_fund_models_cover_all_fin_tables_with_expected_columns() -> None: def test_fund_models_declare_no_foreign_keys_so_no_cascade_writes() -> None: for name, table in Base.metadata.tables.items(): - if not name.startswith("fin_"): + if name not in FUND_TABLES: continue for column in table.columns: assert not column.foreign_keys, f"{name}.{column.name}" diff --git a/tests/unit/service/test_bootstrap.py b/tests/unit/service/test_bootstrap.py index d5905df..75cedb6 100644 --- a/tests/unit/service/test_bootstrap.py +++ b/tests/unit/service/test_bootstrap.py @@ -11,3 +11,7 @@ def test_bootstrap_assembles_common_model_and_tool_services() -> None: assert isinstance(factory._intent_classifier, IntentClassifier) assert factory._intent_endpoint_resolver is not None assert factory._tool_executor.registry.get("check_suitability").read_only is True + knowledge_tool = factory._tool_executor.registry.get("query_knowledge") + assert knowledge_tool.required_permission == "knowledge:query" + assert knowledge_tool.allowed_roles == ("visitor", "customer") + assert knowledge_tool.read_only is True diff --git a/tests/unit/service/test_knowledge_authority.py b/tests/unit/service/test_knowledge_authority.py new file mode 100644 index 0000000..79982f1 --- /dev/null +++ b/tests/unit/service/test_knowledge_authority.py @@ -0,0 +1,61 @@ +import pytest + +from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery +from app.service.knowledge_authority import KnowledgeMysqlAuthority + + +class FakeRow: + id = 101 + title = "申购规则" + content_text = '{"answer":"工作日确认"}' + version = "v2" + milvus_collection = "fin_policy_collection" + + +class FakeSession: + statement = None + + async def scalars(self, statement): + self.statement = statement + return [FakeRow()] + + +@pytest.mark.asyncio +async def test_authority_filters_to_published_active_effective_knowledge() -> None: + session = FakeSession() + authority = KnowledgeMysqlAuthority(session) + + hits = await authority.filter_published(( + KnowledgeHit( + knowledge_id="101", collection="fin_policy_collection", snippet="摘要", score=0.91, + ), + )) + + statement = str(session.statement) + assert "review_status" in statement + assert "status" in statement + assert "effective_date" in statement + assert "expire_date" in statement + assert hits[0].answer == "工作日确认" + assert hits[0].version == "v2" + + +@pytest.mark.asyncio +async def test_authority_keyword_fallback_only_returns_effective_public_knowledge() -> None: + session = FakeSession() + authority = KnowledgeMysqlAuthority(session) + + hits = await authority.search_keyword( + KnowledgeQuery(query="基金 申购确认", intents=("policy_explain",)), + ("fin_policy_collection",), + 5, + ) + + statement = str(session.statement) + assert "milvus_collection" in statement + assert "content_text" in statement + assert "review_status" in statement + assert "effective_date" in statement + assert hits[0].knowledge_id == "101" + assert hits[0].collection == "fin_policy_collection" + assert hits[0].answer == "工作日确认" diff --git a/tests/unit/service/test_knowledge_retrieval.py b/tests/unit/service/test_knowledge_retrieval.py new file mode 100644 index 0000000..8bc112b --- /dev/null +++ b/tests/unit/service/test_knowledge_retrieval.py @@ -0,0 +1,87 @@ +import pytest + +from app.core.contracts import RequestContext +from app.core.errors import RecoverableAgentError +from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery +from app.service.knowledge_config import KnowledgeRuntimeConfig +from app.service.knowledge_retrieval_service import KnowledgeRetrievalService + +# ruff: noqa: E501 + + +class FakeEmbedder: + async def embed(self, text: str) -> list[float]: + assert text == "开户" + return [0.1] * 1024 + + +class FakeVectorStore: + def __init__(self) -> None: + self.calls: list[tuple[str, int]] = [] + + async def search(self, collection: str, vector: list[float], top_k: int) -> list[dict[str, object]]: + assert len(vector) == 1024 + self.calls.append((collection, top_k)) + return [] + + +class FakeAuthority: + async def filter_published(self, hits: tuple[object, ...]) -> list[object]: + return [] + + async def search_keyword(self, query: object, collections: tuple[str, ...], top_k: int) -> list[object]: + return [] + + +class BrokenVectorStore: + async def search(self, collection: str, vector: list[float], top_k: int) -> list[dict[str, object]]: + raise RecoverableAgentError("知识检索不可用") + + +class FallbackAuthority: + def __init__(self) -> None: + self.calls: list[tuple[tuple[str, ...], int]] = [] + + async def filter_published(self, hits: tuple[object, ...]) -> list[object]: + return [] + + async def search_keyword(self, query: KnowledgeQuery, collections: tuple[str, ...], top_k: int) -> list[KnowledgeHit]: + self.calls.append((collections, top_k)) + return [KnowledgeHit( + knowledge_id="101", collection="fin_policy_collection", snippet="确认规则", + answer="工作日确认", score=1.0, + )] + + +@pytest.mark.asyncio +async def test_search_uses_faq_collection_for_faq_only() -> None: + vector_store = FakeVectorStore() + service = KnowledgeRetrievalService( + FakeEmbedder(), vector_store, KnowledgeRuntimeConfig(), FakeAuthority() + ) + + result = await service.search( + KnowledgeQuery(query="开户", intents=("faq",)), + RequestContext(user_id="visitor-1", trace_id="trace", roles=("visitor",), data_scope="public"), + ) + + assert vector_store.calls == [("fin_faq_collection", 3)] + assert result.searched_collections == ("fin_faq_collection",) + + +@pytest.mark.asyncio +async def test_milvus_failure_falls_back_to_published_active_unexpired_knowledge() -> None: + authority = FallbackAuthority() + service = KnowledgeRetrievalService( + FakeEmbedder(), BrokenVectorStore(), KnowledgeRuntimeConfig(), authority + ) + + result = await service.search( + KnowledgeQuery(query="开户", intents=("policy_explain",)), + RequestContext(user_id="visitor-1", trace_id="trace", roles=("visitor",), data_scope="public"), + ) + + assert authority.calls == [(("fin_policy_collection",), 5)] + assert result.degraded is True + assert result.degradation_reason == "milvus_unavailable" + assert result.hits[0].answer == "工作日确认" diff --git a/tests/unit/service/test_knowledge_tool_service.py b/tests/unit/service/test_knowledge_tool_service.py new file mode 100644 index 0000000..1310736 --- /dev/null +++ b/tests/unit/service/test_knowledge_tool_service.py @@ -0,0 +1,131 @@ +import pytest + +from app.core.contracts import RequestContext +from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery, KnowledgeSearchResult +from app.service import knowledge_tool_service +from app.service.knowledge_tool_service import DatabaseEmbeddingAdapter, query_knowledge_tool + + +class FakeGateway: + def __init__(self) -> None: + self.calls: list[tuple[str, str, int]] = [] + + async def embed(self, *, endpoint_code: str, text: str, timeout_ms: int) -> list[float]: + self.calls.append((endpoint_code, text, timeout_ms)) + return [0.1] * 1024 + + +@pytest.mark.asyncio +async def test_embedding_adapter_uses_single_text_gateway_contract() -> None: + gateway = FakeGateway() + adapter = DatabaseEmbeddingAdapter("knowledge-embedding", 15000, gateway=gateway) + + vector = await adapter.embed("基金开户") + + assert len(vector) == 1024 + assert gateway.calls == [("knowledge-embedding", "基金开户", 15000)] + + +@pytest.mark.asyncio +async def test_query_tool_degrades_when_embedding_endpoint_is_unconfigured(monkeypatch) -> None: + class Settings: + knowledge_embedding_endpoint_code = "" + + monkeypatch.setattr("app.service.knowledge_tool_service.get_settings", lambda: Settings()) + + result = await query_knowledge_tool( + KnowledgeQuery(query="基金开户", intents=("faq",)), + RequestContext( + user_id="visitor-1", trace_id="trace", roles=("visitor",), data_scope="public" + ), + ) + + assert result.degraded is True + assert result.degradation_reason == "embedding_endpoint_unconfigured" + + +@pytest.mark.asyncio +async def test_query_tool_uses_configured_embedding_endpoint_and_read_only_dependencies( + monkeypatch, +) -> None: + class Settings: + knowledge_embedding_endpoint_code = "knowledge-embedding" + knowledge_embedding_timeout_ms = 15000 + milvus_uri = "http://milvus:19530" + milvus_token = "" + + class FakeGateway: + calls: list[tuple[str, str, int]] = [] + + async def embed( + self, *, endpoint_code: str, text: str, timeout_ms: int + ) -> list[float]: + self.calls.append((endpoint_code, text, timeout_ms)) + return [0.1] * 1024 + + class FakeMilvus: + def __init__(self, uri: str, token: str | None) -> None: + self.uri = uri + self.token = token + + class FakeSession: + async def __aenter__(self) -> object: + return object() + + async def __aexit__(self, exc_type, exc, traceback) -> None: + return None + + class FakeAuthority: + def __init__(self, session: object) -> None: + self.session = session + + class FakeRetrievalService: + def __init__(self, embedder, vector_store, config, authority) -> None: + self.embedder = embedder + self.vector_store = vector_store + self.config = config + self.authority = authority + + async def search( + self, query: KnowledgeQuery, context: RequestContext + ) -> KnowledgeSearchResult: + vector = await self.embedder.embed(query.query) + assert len(vector) == 1024 + assert isinstance(self.vector_store, FakeMilvus) + assert isinstance(self.authority, FakeAuthority) + assert self.config.routes["faq"] == ("fin_faq_collection", 3) + assert context.data_scope == "public" + return KnowledgeSearchResult( + hits=( + KnowledgeHit( + knowledge_id="1", + collection="fin_faq_collection", + snippet="snippet", + answer="answer", + ), + ), + searched_collections=("fin_faq_collection",), + ) + + gateway = FakeGateway() + monkeypatch.setattr(knowledge_tool_service, "get_settings", lambda: Settings()) + monkeypatch.setattr(knowledge_tool_service, "DatabaseModelGateway", lambda: gateway) + monkeypatch.setattr(knowledge_tool_service, "MilvusKnowledgeClient", FakeMilvus) + monkeypatch.setattr(knowledge_tool_service, "KnowledgeMysqlAuthority", FakeAuthority) + monkeypatch.setattr( + knowledge_tool_service, "KnowledgeRetrievalService", FakeRetrievalService + ) + monkeypatch.setattr(knowledge_tool_service, "SessionFactory", FakeSession) + + result = await query_knowledge_tool( + KnowledgeQuery(query="基金开户", intents=("faq",)), + RequestContext( + user_id="visitor-1", + trace_id="trace", + roles=("visitor",), + data_scope="public", + ), + ) + + assert result.hits[0].answer == "answer" + assert gateway.calls == [("knowledge-embedding", "基金开户", 15000)] diff --git a/tests/unit/worker/test_runtime_worker_dispatch.py b/tests/unit/worker/test_runtime_worker_dispatch.py index 4d3beef..4c9f502 100644 --- a/tests/unit/worker/test_runtime_worker_dispatch.py +++ b/tests/unit/worker/test_runtime_worker_dispatch.py @@ -34,7 +34,7 @@ async def test_worker_restores_visitor_without_identity_repository_call() -> Non ) assert context.roles == ("visitor",) - assert context.permissions == ("agent:run",) + assert context.permissions == ("agent:run", "knowledge:query") assert context.data_scope == "public" runtime.resolve_identity.assert_not_awaited()