feat: add governed public knowledge retrieval
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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}"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 == "工作日确认"
|
||||
@@ -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 == "工作日确认"
|
||||
@@ -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)]
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user