diff --git a/python/app/rag/embeddings/reranker.py b/python/app/rag/embeddings/reranker.py index 552b65d..5d8c668 100755 --- a/python/app/rag/embeddings/reranker.py +++ b/python/app/rag/embeddings/reranker.py @@ -13,20 +13,24 @@ def normalize_text(text: str) -> str: return _WHITESPACE_RE.sub(" ", text).strip().lower() -def document_signature(document: Document) -> str: - """Generate a signature for a document based on its content.""" +def document_signature(document: Document) -> bytes: + """Generate a compact signature for a document's normalized content.""" normalized_content = normalize_text(document.page_content) - return sha1(normalized_content.encode("utf-8")).hexdigest() + return sha1(normalized_content.encode("utf-8")).digest() -def deduplicate_documents(documents: list[Document]) -> list[Document]: +def deduplicate_documents( + documents: list[Document], + seen_signatures: set[bytes] | None = None, + max_seen_signatures: int | None = None, +) -> list[Document]: """ - Remove duplicate documents based on content similarity. + Remove duplicate documents by normalized content. - Uses SHA1 hashing of normalized content to detect duplicates. - Safe for low-memory environments (no ML models). + A shared signature set deduplicates across batches. Its optional cap bounds + extra memory; duplicates remain deduplicated within each individual batch. """ - seen_signatures: set[str] = set() + batch_signatures: set[bytes] = set() deduplicated: list[Document] = [] for document in documents: @@ -34,10 +38,20 @@ def deduplicate_documents(documents: list[Document]) -> list[Document]: continue signature = document_signature(document) - if signature in seen_signatures: + if signature in batch_signatures: continue + batch_signatures.add(signature) - seen_signatures.add(signature) + if seen_signatures is not None and signature in seen_signatures: + continue + if ( + seen_signatures is not None + and ( + max_seen_signatures is None + or len(seen_signatures) < max_seen_signatures + ) + ): + seen_signatures.add(signature) deduplicated.append(document) return deduplicated diff --git a/python/app/rag/embeddings/vector_db.py b/python/app/rag/embeddings/vector_db.py index 5589baa..2c3330e 100755 --- a/python/app/rag/embeddings/vector_db.py +++ b/python/app/rag/embeddings/vector_db.py @@ -20,8 +20,16 @@ def __init__(self, collection_name): ) def add_documents(self, documents, batch_size=100): + # Keep one content-signature set across batches without copying the + # complete document corpus into a second list. + max_seen_signatures = int(os.getenv("RAG_DEDUP_MAX_SIGNATURES", "50000")) + seen_signatures: set[bytes] = set() for i in range(0, len(documents), batch_size): - batch = deduplicate_documents(documents[i : i + batch_size]) + batch = deduplicate_documents( + documents[i : i + batch_size], + seen_signatures=seen_signatures, + max_seen_signatures=max_seen_signatures, + ) if not batch: continue diff --git a/python/tests/test_vector_db.py b/python/tests/test_vector_db.py new file mode 100644 index 0000000..c562346 --- /dev/null +++ b/python/tests/test_vector_db.py @@ -0,0 +1,59 @@ +import os +import unittest +from unittest.mock import MagicMock, patch + +from langchain_core.documents import Document + +from app.rag.embeddings.vector_db import vector_db + + +class VectorDatabaseIngestionTests(unittest.TestCase): + def test_duplicate_chunks_across_batches_are_submitted_once(self): + unique = [ + Document( + page_content=f"Source passage {index}", + metadata={"job_id": "job-1", "chunk_index": index}, + ) + for index in range(100) + ] + duplicated = [ + Document( + page_content=f" source passage {index} ", + metadata={"job_id": "job-1", "chunk_index": index + 100}, + ) + for index in range(100) + ] + collection = MagicMock() + database = vector_db.__new__(vector_db) + database.collection = collection + + database.add_documents(unique + duplicated, batch_size=100) + + submitted = sum( + len(call.args[0]) + for call in collection.add_documents.call_args_list + ) + self.assertEqual(submitted, 100) + self.assertEqual(collection.add_documents.call_count, 1) + + + def test_signature_memory_cap_is_respected_across_batches(self): + documents = [ + Document(page_content=content, metadata={"job_id": "job-1"}) + for content in ("alpha", "beta", "gamma", "alpha", "beta", "gamma") + ] + collection = MagicMock() + database = vector_db.__new__(vector_db) + database.collection = collection + + with patch.dict(os.environ, {"RAG_DEDUP_MAX_SIGNATURES": "2"}): + database.add_documents(documents, batch_size=3) + + submitted_batches = [ + len(call.args[0]) + for call in collection.add_documents.call_args_list + ] + self.assertEqual(submitted_batches, [3, 1]) + +if __name__ == "__main__": + unittest.main()