Files
Mutual_Fund/tests/test_ingestion.py
T

86 lines
2.9 KiB
Python

import unittest
from unittest.mock import AsyncMock
from rag.embedding import EMBEDDING_DIMENSION
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] * EMBEDDING_DIMENSION])
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] * EMBEDDING_DIMENSION])
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] * EMBEDDING_DIMENSION, [0.0] * EMBEDDING_DIMENSION])
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()