test: queue manager
This commit is contained in:
parent
4328fb3d2a
commit
faefb9ac86
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue