Files
group_fqcd_jr/tests/unit/service/test_knowledge_publication_service.py
T

120 lines
4.1 KiB
Python
Raw Normal View History

from dataclasses import dataclass
import pytest
from app.service.knowledge_publication_service import (
KnowledgePublicationError,
KnowledgePublicationService,
)
@dataclass(frozen=True)
class Record:
qa_id: str
milvus_collection: str
retrieval_text: str
title: str = "基金开户"
snippet: str = "基金开户"
tags: tuple[str, ...] = ("开户",)
version: str = "v5.8"
class FakeEmbedder:
def __init__(self, vector: list[float] | None = None) -> None:
self.vector = vector or [0.1] * 1024
self.calls: list[str] = []
async def embed(self, text: str) -> list[float]:
self.calls.append(text)
return self.vector
class FakeStore:
def __init__(self) -> None:
self.staged: list[Record] = []
self.published: list[int] = []
self.disabled: list[int] = []
async def stage(self, records: tuple[Record, ...]) -> dict[str, int]:
self.staged.extend(records)
return {record.qa_id: index for index, record in enumerate(records, start=101)}
async def publish(self, knowledge_ids: tuple[int, ...], reviewer_id: int) -> None:
assert reviewer_id == 9
self.published.extend(knowledge_ids)
async def disable(self, knowledge_ids: tuple[int, ...]) -> None:
self.disabled.extend(knowledge_ids)
class FakeVectors:
def __init__(self, *, fail_collection: str | None = None) -> None:
self.fail_collection = fail_collection
self.upserts: list[tuple[str, tuple[dict[str, object], ...]]] = []
self.deleted: list[tuple[str, tuple[str, ...]]] = []
async def upsert(self, collection: str, records: tuple[dict[str, object], ...]) -> None:
self.upserts.append((collection, records))
if collection == self.fail_collection:
raise RuntimeError("milvus unavailable")
async def delete(self, collection: str, knowledge_ids: tuple[str, ...]) -> None:
self.deleted.append((collection, knowledge_ids))
@pytest.mark.asyncio
async def test_publication_stages_vectors_then_publishes_after_all_collections_succeed() -> None:
store = FakeStore()
vectors = FakeVectors()
service = KnowledgePublicationService(FakeEmbedder(), store, vectors)
records = (
Record("FAQ-001", "fin_faq_collection", "标准问题:基金开户"),
Record("POL-001", "fin_policy_collection", "标准问题:风险测评"),
)
result = await service.publish(records, reviewer_id=9)
assert result.knowledge_ids == {"FAQ-001": 101, "POL-001": 102}
assert [collection for collection, _payload in vectors.upserts] == [
"fin_faq_collection", "fin_policy_collection"
]
assert store.published == [101, 102]
assert store.disabled == []
payload = vectors.upserts[0][1][0]
assert payload["knowledge_id"] == "101"
assert len(payload["embedding"]) == 1024
@pytest.mark.asyncio
async def test_vector_failure_keeps_staged_rows_unpublished_and_compensates_vectors() -> None:
store = FakeStore()
vectors = FakeVectors(fail_collection="fin_policy_collection")
service = KnowledgePublicationService(FakeEmbedder(), store, vectors)
records = (
Record("FAQ-001", "fin_faq_collection", "标准问题:基金开户"),
Record("POL-001", "fin_policy_collection", "标准问题:风险测评"),
)
with pytest.raises(KnowledgePublicationError, match="向量写入失败"):
await service.publish(records, reviewer_id=9)
assert store.published == []
assert store.disabled == [101, 102]
assert vectors.deleted == [
("fin_faq_collection", ("101",)),
("fin_policy_collection", ("102",)),
]
@pytest.mark.asyncio
async def test_invalid_embedding_dimension_blocks_all_database_and_vector_writes() -> None:
store = FakeStore()
vectors = FakeVectors()
service = KnowledgePublicationService(FakeEmbedder([0.1] * 512), store, vectors)
with pytest.raises(KnowledgePublicationError, match="1024"):
await service.publish((Record("FAQ-001", "fin_faq_collection", "基金开户"),), reviewer_id=9)
assert store.staged == []
assert vectors.upserts == []