diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index d8d89bae..788c36a8 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -236,3 +236,222 @@ class TestQueueProcessing: assert "None" in summary_work_unit_key assert "summary" in summary_work_unit_key assert "workspace1" in summary_work_unit_key + + @pytest.mark.asyncio + async def test_representation_batching_respects_token_limits( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + ) -> None: + """Test that representation tasks are batched based on token limits""" + from unittest.mock import patch + + session, peers = sample_session_with_peers + peer = peers[0] + + # Create messages with token counts that exceed batch limit + # Total: 2000 + 2000 + 3000 = 7000 tokens (> 4096 limit defined in settings) + messages = [ + models.Message( + id=i, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content=f"Test message {i}", + token_count=token_count, + ) + for i, token_count in enumerate([2000, 2000, 3000]) + ] + + # Create queue items with token counts + payloads = [ + create_queue_payload( # type: ignore[reportUnknownArgumentType] + message=msg, + task_type="representation", + sender_name=peer.name, + target_name=peer.name, + ) + for msg in messages + ] + + # Add items with token counts + from src.deriver.utils import get_work_unit_key + + queue_items: list[models.QueueItem] = [] + for payload, message in zip(payloads, messages, strict=False): + task_type = payload.get("task_type", "unknown") + work_unit_key = get_work_unit_key(task_type, payload) + + queue_item = models.QueueItem( + session_id=session.id, + task_type=task_type, + work_unit_key=work_unit_key, + payload=payload, + processed=False, + token_count=message.token_count, + ) + db_session.add(queue_item) + queue_items.append(queue_item) + + await db_session.commit() + for item in queue_items: + await db_session.refresh(item) + + # Mock process_items to capture batches + processed_batches: list[dict[str, Any]] = [] + + async def mock_process_items( + task_type: str, queue_payloads: list[dict[str, Any]] + ) -> None: + processed_batches.append( + { + "task_type": task_type, + "payload_count": len(queue_payloads), + } + ) + + # Process work unit and verify batching + qm = QueueManager() + with patch( + "src.deriver.queue_manager.process_items", side_effect=mock_process_items + ): + await qm.process_work_unit(queue_items[0].work_unit_key) + + # Should create 2 batches due to token limits + assert len(processed_batches) == 2 + assert processed_batches[0]["payload_count"] == 2 # 2000 + 2000 + assert processed_batches[1]["payload_count"] == 1 # 3000 + + @pytest.mark.asyncio + async def test_hard_batch_size_limit( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + ) -> None: + """Test that get_message_batch respects the hard limit of 10 messages""" + session, peers = sample_session_with_peers + peer = peers[0] + + # Create 15 messages to test batch size limit + messages = [ + models.Message( + id=i, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content=f"Test message {i}", + token_count=100, # Small tokens to avoid token-based batching + ) + for i in range(15) + ] + + payloads = [ + create_queue_payload(msg, "representation", peer.name, peer.name) + for msg in messages + ] + + # Create queue items + from src.deriver.utils import get_work_unit_key + + queue_items: list[models.QueueItem] = [] + for payload, message in zip(payloads, messages, strict=False): + task_type = payload.get("task_type", "unknown") + work_unit_key = get_work_unit_key(task_type, payload) + + queue_item = models.QueueItem( + session_id=session.id, + task_type=task_type, + work_unit_key=work_unit_key, + payload=payload, + processed=False, + token_count=message.token_count, + ) + db_session.add(queue_item) + queue_items.append(queue_item) + + await db_session.commit() + + # Test batch size limit + qm = QueueManager() + batch = await qm.get_message_batch(queue_items[0].work_unit_key, limit=10) + assert len(batch) == 10 # Should respect hard limit + + @pytest.mark.asyncio + async def test_single_message_processing( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + ) -> None: + """Test that multiple summary messages in same work unit are processed separately""" + from unittest.mock import patch + + session, peers = sample_session_with_peers + peer = peers[0] + + # Create two summary messages + messages = [ + models.Message( + id=999, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content="First summary message", + token_count=500, + ), + models.Message( + id=1000, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content="Second summary message", + token_count=600, + ), + ] + + # Create payloads and queue items + queue_items: list[models.QueueItem] = [] + for i, message in enumerate(messages): + payload = create_queue_payload( + message, "summary", message_seq_in_session=i + 1 + ) + from src.deriver.utils import get_work_unit_key + + work_unit_key = get_work_unit_key("summary", payload) + + queue_item = models.QueueItem( + session_id=session.id, + task_type="summary", + work_unit_key=work_unit_key, + payload=payload, + processed=False, + token_count=message.token_count, + ) + db_session.add(queue_item) + queue_items.append(queue_item) + + await db_session.commit() + + # Mock and process work unit + processed_batches: list[dict[str, Any]] = [] + + async def mock_process_items( + task_type: str, queue_payloads: list[dict[str, Any]] + ) -> None: + processed_batches.append( + {"task_type": task_type, "payload_count": len(queue_payloads)} + ) + + qm = QueueManager() + work_unit_key = queue_items[0].work_unit_key + with patch( + "src.deriver.queue_manager.process_items", side_effect=mock_process_items + ): + await qm.process_work_unit(work_unit_key) + + # Verify both messages were processed in separate batches + assert len(processed_batches) == 2 + assert all(batch["task_type"] == "summary" for batch in processed_batches) + assert all(batch["payload_count"] == 1 for batch in processed_batches)