120 lines
4.1 KiB
Python
120 lines
4.1 KiB
Python
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 == []
|