| """ |
| Unit tests for PGKVStorage.upsert batch optimization (PR #2742 fixes). |
| |
| Verifies: |
| 1. Each namespace builds correct tuple ordering matching SQL positional params. |
| 2. _run_with_retry is used (not the removed PostgreSQLDB.executemany wrapper). |
| 3. Sub-batching splits data when len(data) > _max_batch_size. |
| 4. Unknown namespace raises ValueError. |
| 5. Empty data returns without any DB call. |
| """ |
|
|
| import json |
| import pytest |
| import numpy as np |
| from unittest.mock import AsyncMock, MagicMock |
| from lightrag.kg.postgres_impl import PGDocStatusStorage, PGKVStorage, PGVectorStorage |
| from lightrag.namespace import NameSpace |
| from lightrag.utils import EmbeddingFunc |
|
|
|
|
| |
| |
| |
|
|
| GLOBAL_CONFIG = {"embedding_batch_num": 10} |
|
|
|
|
| def make_storage(namespace: str) -> PGKVStorage: |
| """Construct a PGKVStorage instance with a mocked db.""" |
| db = MagicMock() |
| captured: list[tuple] = [] |
| retry_kwargs: list[dict] = [] |
|
|
| async def fake_run_with_retry(operation, **kwargs): |
| """Call the closure with a mock connection to capture executemany args.""" |
| retry_kwargs.append(kwargs) |
| mock_conn = AsyncMock() |
| await operation(mock_conn) |
| |
| for call in mock_conn.executemany.call_args_list: |
| captured.append((call.args[0], call.args[1])) |
|
|
| db._run_with_retry = AsyncMock(side_effect=fake_run_with_retry) |
| db.workspace = "test_ws" |
|
|
| storage = PGKVStorage.__new__(PGKVStorage) |
| storage.namespace = namespace |
| storage.workspace = "test_ws" |
| storage.global_config = GLOBAL_CONFIG |
| storage.db = db |
| storage.__post_init__() |
|
|
| storage._captured = captured |
| storage._retry_kwargs = retry_kwargs |
| return storage |
|
|
|
|
| def make_doc_status_storage() -> PGDocStatusStorage: |
| """Construct a PGDocStatusStorage instance with a mocked db.""" |
| db = MagicMock() |
| captured: list[tuple] = [] |
| retry_kwargs: list[dict] = [] |
|
|
| async def fake_run_with_retry(operation, **kwargs): |
| retry_kwargs.append(kwargs) |
| mock_conn = AsyncMock() |
| tx = AsyncMock() |
| tx.__aenter__.return_value = tx |
| tx.__aexit__.return_value = False |
| mock_conn.transaction = MagicMock(return_value=tx) |
| await operation(mock_conn) |
| for call in mock_conn.executemany.call_args_list: |
| captured.append((call.args[0], call.args[1])) |
|
|
| db._run_with_retry = AsyncMock(side_effect=fake_run_with_retry) |
| db.workspace = "test_ws" |
|
|
| storage = PGDocStatusStorage.__new__(PGDocStatusStorage) |
| storage.namespace = NameSpace.DOC_STATUS |
| storage.workspace = "test_ws" |
| storage.global_config = GLOBAL_CONFIG |
| storage.db = db |
| storage._captured = captured |
| storage._retry_kwargs = retry_kwargs |
| return storage |
|
|
|
|
| def make_vector_storage(namespace: str) -> PGVectorStorage: |
| """Construct a PGVectorStorage instance with a mocked db and embedding func.""" |
| db = MagicMock() |
| captured: list[tuple] = [] |
| retry_kwargs: list[dict] = [] |
|
|
| async def fake_run_with_retry(operation, **kwargs): |
| retry_kwargs.append(kwargs) |
| mock_conn = AsyncMock() |
| await operation(mock_conn) |
| for call in mock_conn.executemany.call_args_list: |
| captured.append((call.args[0], call.args[1])) |
|
|
| db._run_with_retry = AsyncMock(side_effect=fake_run_with_retry) |
| db.workspace = "test_ws" |
|
|
| async def embed_func(texts, **kwargs): |
| return np.array([[0.1, 0.2, 0.3] for _ in texts], dtype=np.float32) |
|
|
| embedding = EmbeddingFunc( |
| embedding_dim=3, |
| func=embed_func, |
| model_name="test_model", |
| ) |
| storage = PGVectorStorage( |
| namespace=namespace, |
| workspace="test_ws", |
| global_config={ |
| "embedding_batch_num": 10, |
| "vector_db_storage_cls_kwargs": { |
| "cosine_better_than_threshold": 0.5, |
| }, |
| }, |
| embedding_func=embedding, |
| ) |
| storage.db = db |
| storage._captured = captured |
| storage._retry_kwargs = retry_kwargs |
| return storage |
|
|
|
|
| |
| |
| |
|
|
|
|
| def test_max_batch_size_is_constant(): |
| storage = make_storage(NameSpace.KV_STORE_TEXT_CHUNKS) |
| assert storage._max_batch_size == 200 |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_text_chunks_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_TEXT_CHUNKS) |
| data = { |
| "chunk-1": { |
| "tokens": 42, |
| "chunk_order_index": 0, |
| "full_doc_id": "doc-1", |
| "content": "hello world", |
| "file_path": "/a/b.txt", |
| "llm_cache_list": ["cache-key"], |
| } |
| } |
| await storage.upsert(data) |
|
|
| assert len(storage._captured) == 1 |
| sql, rows = storage._captured[0] |
| assert "LIGHTRAG_DOC_CHUNKS" in sql |
| assert len(rows) == 1 |
| row = rows[0] |
| |
| |
| assert row[0] == "test_ws" |
| assert row[1] == "chunk-1" |
| assert row[2] == 42 |
| assert row[3] == 0 |
| assert row[4] == "doc-1" |
| assert row[5] == "hello world" |
| assert row[6] == "/a/b.txt" |
| assert json.loads(row[7]) == ["cache-key"] |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_full_docs_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| data = {"doc-1": {"content": "full text", "file_path": "/path/doc.pdf"}} |
| await storage.upsert(data) |
|
|
| assert len(storage._captured) == 1 |
| _, rows = storage._captured[0] |
| row = rows[0] |
| |
| assert row[0] == "doc-1" |
| assert row[1] == "full text" |
| assert row[2] == "/path/doc.pdf" |
| assert row[3] == "test_ws" |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_llm_cache_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_LLM_RESPONSE_CACHE) |
| data = { |
| "key-1": { |
| "original_prompt": "what is X?", |
| "return": "X is Y", |
| "chunk_id": "chunk-1", |
| "cache_type": "query", |
| "queryparam": {"mode": "hybrid"}, |
| } |
| } |
| await storage.upsert(data) |
|
|
| assert len(storage._captured) == 1 |
| _, rows = storage._captured[0] |
| row = rows[0] |
| |
| assert row[0] == "test_ws" |
| assert row[1] == "key-1" |
| assert row[2] == "what is X?" |
| assert row[3] == "X is Y" |
| assert row[4] == "chunk-1" |
| assert row[5] == "query" |
| assert json.loads(row[6]) == {"mode": "hybrid"} |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_llm_cache_null_queryparam(): |
| storage = make_storage(NameSpace.KV_STORE_LLM_RESPONSE_CACHE) |
| data = { |
| "key-2": { |
| "original_prompt": "prompt", |
| "return": "answer", |
| "cache_type": "extract", |
| } |
| } |
| await storage.upsert(data) |
| _, rows = storage._captured[0] |
| assert rows[0][6] is None |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_full_entities_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_ENTITIES) |
| data = {"ent-1": {"entity_names": ["EntityA", "EntityB"], "count": 2}} |
| await storage.upsert(data) |
|
|
| _, rows = storage._captured[0] |
| row = rows[0] |
| |
| assert row[0] == "test_ws" |
| assert row[1] == "ent-1" |
| assert json.loads(row[2]) == ["EntityA", "EntityB"] |
| assert row[3] == 2 |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_full_relations_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_RELATIONS) |
| data = {"rel-1": {"relation_pairs": [["A", "B"]], "count": 1}} |
| await storage.upsert(data) |
|
|
| _, rows = storage._captured[0] |
| row = rows[0] |
| |
| assert row[0] == "test_ws" |
| assert row[1] == "rel-1" |
| assert json.loads(row[2]) == [["A", "B"]] |
| assert row[3] == 1 |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_entity_chunks_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_ENTITY_CHUNKS) |
| data = {"ec-1": {"chunk_ids": ["c1", "c2"], "count": 2}} |
| await storage.upsert(data) |
|
|
| _, rows = storage._captured[0] |
| row = rows[0] |
| |
| assert row[0] == "test_ws" |
| assert row[1] == "ec-1" |
| assert json.loads(row[2]) == ["c1", "c2"] |
| assert row[3] == 2 |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_relation_chunks_tuple_order(): |
| storage = make_storage(NameSpace.KV_STORE_RELATION_CHUNKS) |
| data = {"rc-1": {"chunk_ids": ["c3"], "count": 1}} |
| await storage.upsert(data) |
|
|
| _, rows = storage._captured[0] |
| row = rows[0] |
| assert row[0] == "test_ws" |
| assert row[1] == "rc-1" |
| assert json.loads(row[2]) == ["c3"] |
| assert row[3] == 1 |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_sub_batching_splits_correctly(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| storage._max_batch_size = 3 |
|
|
| data = {f"doc-{i}": {"content": f"text {i}", "file_path": ""} for i in range(7)} |
| await storage.upsert(data) |
|
|
| |
| assert len(storage._captured) == 3 |
| assert len(storage._captured[0][1]) == 3 |
| assert len(storage._captured[1][1]) == 3 |
| assert len(storage._captured[2][1]) == 1 |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_sub_batching_exact_multiple(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| storage._max_batch_size = 3 |
|
|
| data = {f"doc-{i}": {"content": f"text {i}", "file_path": ""} for i in range(6)} |
| await storage.upsert(data) |
|
|
| |
| assert len(storage._captured) == 2 |
| assert len(storage._captured[0][1]) == 3 |
| assert len(storage._captured[1][1]) == 3 |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_empty_data_no_db_call(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| await storage.upsert({}) |
| assert len(storage._captured) == 0 |
| storage.db._run_with_retry.assert_not_called() |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_upsert_unknown_namespace_raises(): |
| storage = make_storage("unknown_namespace") |
| with pytest.raises(ValueError, match="Unknown namespace"): |
| await storage.upsert({"k": {"v": 1}}) |
|
|
|
|
| |
| |
| |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_multiple_records_single_batch(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| data = { |
| "doc-1": {"content": "text 1", "file_path": "/a"}, |
| "doc-2": {"content": "text 2", "file_path": "/b"}, |
| "doc-3": {"content": "text 3", "file_path": "/c"}, |
| } |
| await storage.upsert(data) |
|
|
| |
| assert len(storage._captured) == 1 |
| _, rows = storage._captured[0] |
| assert len(rows) == 3 |
| ids = {row[0] for row in rows} |
| assert ids == {"doc-1", "doc-2", "doc-3"} |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_kv_upsert_passes_timing_label(): |
| storage = make_storage(NameSpace.KV_STORE_FULL_DOCS) |
| await storage.upsert({"doc-1": {"content": "text 1", "file_path": "/a"}}) |
|
|
| assert storage._retry_kwargs[0]["timing_label"] == ( |
| f"test_ws PGKVStorage.upsert[{NameSpace.KV_STORE_FULL_DOCS}]" |
| ) |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_doc_status_upsert_passes_timing_label(): |
| storage = make_doc_status_storage() |
| await storage.upsert( |
| { |
| "doc-1": { |
| "content_summary": "summary", |
| "content_length": 12, |
| "chunks_count": 1, |
| "status": "processed", |
| "file_path": "/a.txt", |
| "chunks_list": ["chunk-1"], |
| "metadata": {"source": "test"}, |
| "created_at": "2024-01-01T00:00:00+00:00", |
| "updated_at": "2024-01-01T00:00:00+00:00", |
| } |
| } |
| ) |
|
|
| assert storage._retry_kwargs[0]["timing_label"] == ( |
| "test_ws PGDocStatusStorage.upsert" |
| ) |
|
|
|
|
| @pytest.mark.asyncio |
| async def test_vector_upsert_passes_timing_label(): |
| storage = make_vector_storage(NameSpace.VECTOR_STORE_CHUNKS) |
| await storage.upsert( |
| { |
| "chunk-1": { |
| "tokens": 42, |
| "chunk_order_index": 0, |
| "full_doc_id": "doc-1", |
| "content": "hello world", |
| "file_path": "/a/b.txt", |
| } |
| } |
| ) |
|
|
| assert storage._retry_kwargs[0]["timing_label"] == ( |
| f"test_ws PGVectorStorage.upsert[{NameSpace.VECTOR_STORE_CHUNKS}]" |
| ) |
|
|