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()