85 lines
2.8 KiB
Python
85 lines
2.8 KiB
Python
import unittest
|
|||
|
|
from unittest.mock import AsyncMock
|
||
|
|
|
||
|
|
from rag.ingestion import ingest_document_atomic
|
||
|
|
|
||
|
|
|
||
|
|
class AtomicIngestionTests(unittest.IsolatedAsyncioTestCase):
|
||
|
|
async def test_validation_completes_before_embedding(self):
|
||
|
|
milvus = AsyncMock()
|
||
|
|
embedder = AsyncMock()
|
||
|
|
|
||
|
|
with self.assertRaises(ValueError):
|
||
|
|
await ingest_document_atomic(
|
||
|
|
"Q: only a question", "doc-1", "FAQ", "fin_faq", "qa_pair",
|
||
|
|
milvus_client=milvus, embedder=embedder,
|
||
|
|
)
|
||
|
|
|
||
|
|
embedder.assert_not_awaited()
|
||
|
|
milvus.insert.assert_not_awaited()
|
||
|
|
|
||
|
|
async def test_embedding_failure_cleans_up_document_rows(self):
|
||
|
|
milvus = AsyncMock()
|
||
|
|
embedder = AsyncMock(side_effect=RuntimeError("embedding down"))
|
||
|
|
|
||
|
|
with self.assertRaises(RuntimeError):
|
||
|
|
await ingest_document_atomic(
|
||
|
|
"plain text", "doc-2", "Policy", "fin_policy", "default",
|
||
|
|
milvus_client=milvus, embedder=embedder,
|
||
|
|
)
|
||
|
|
|
||
|
|
milvus.delete.assert_awaited_once_with(
|
||
|
|
collection_name="fin_policy", filter='doc_id == "doc-2"'
|
||
|
|
)
|
||
|
|
|
||
|
|
async def test_milvus_failure_cleans_up_partial_document_rows(self):
|
||
|
|
milvus = AsyncMock()
|
||
|
|
milvus.insert.side_effect = RuntimeError("insert down")
|
||
|
|
embedder = AsyncMock(return_value=[[0.0] * 768])
|
||
|
|
|
||
|
|
with self.assertRaises(RuntimeError):
|
||
|
|
await ingest_document_atomic(
|
||
|
|
"plain text", "doc-3", "Policy", "fin_policy", "default",
|
||
|
|
milvus_client=milvus, embedder=embedder,
|
||
|
|
)
|
||
|
|
|
||
|
|
milvus.delete.assert_awaited_once_with(
|
||
|
|
collection_name="fin_policy", filter='doc_id == "doc-3"'
|
||
|
|
)
|
||
|
|
|
||
|
|
async def test_success_inserts_complete_document_rows(self):
|
||
|
|
milvus = AsyncMock()
|
||
|
|
embedder = AsyncMock(return_value=[[0.0] * 768])
|
||
|
|
|
||
|
|
result = await ingest_document_atomic(
|
||
|
|
"plain text", "doc-4", "Policy", "fin_policy", "default",
|
||
|
|
milvus_client=milvus, embedder=embedder,
|
||
|
|
)
|
||
|
|
|
||
|
|
self.assertEqual(result["doc_id"], "doc-4")
|
||
|
|
milvus.delete.assert_not_awaited()
|
||
|
|
milvus.insert.assert_awaited_once()
|
||
|
|
|
||
|
|
async def test_cleans_text_before_embedding_and_milvus_insert(self):
|
||
|
|
milvus = AsyncMock()
|
||
|
|
embedder = AsyncMock(return_value=[[0.0] * 768, [0.0] * 768])
|
||
|
|
|
||
|
|
await ingest_document_atomic(
|
||
|
|
"\ufeff基金名称:示例基金\r\n\r\n\r\n风险等级:R3\t ",
|
||
|
|
"doc-clean",
|
||
|
|
"示例基金",
|
||
|
|
"fin_fund_doc",
|
||
|
|
"default",
|
||
|
|
milvus_client=milvus,
|
||
|
|
embedder=embedder,
|
||
|
|
)
|
||
|
|
|
||
|
|
self.assertEqual(
|
||
|
|
embedder.await_args.args[0],
|
||
|
|
["基金名称:示例基金", "风险等级:R3"],
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|