test: queue manager

This commit is contained in:
Rajat Ahuja 2025-09-05 16:12:35 -04:00
parent 4328fb3d2a
commit faefb9ac86
1 changed files with 219 additions and 0 deletions

View File

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