From e6c580b660675ad11431b67f3a8773830c03f0d1 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 12:18:03 -0400 Subject: [PATCH] refactor: queue manager --- src/config.py | 8 ++ src/deriver/__init__.py | 4 +- src/deriver/consumer.py | 44 ++++++--- src/deriver/deriver.py | 10 ++ src/deriver/queue_manager.py | 124 ++++++++++++++++--------- src/deriver/queue_payload.py | 6 ++ tests/deriver/test_queue_processing.py | 6 +- 7 files changed, 140 insertions(+), 62 deletions(-) diff --git a/src/config.py b/src/config.py index 68e9039d..02151cc8 100644 --- a/src/config.py +++ b/src/config.py @@ -212,6 +212,14 @@ class DeriverSettings(HonchoSettings): int, Field(default=100, gt=0, le=500) ] = 100 + REPRESENTATION_BATCH_MAX_TOKENS: Annotated[ + int, + Field( + default=4096, + ge=1, + ), + ] = 4096 + class DialecticSettings(HonchoSettings): model_config = SettingsConfigDict(env_prefix="DIALECTIC_", extra="ignore") # pyright: ignore diff --git a/src/deriver/__init__.py b/src/deriver/__init__.py index 26d3733d..707ba77d 100644 --- a/src/deriver/__init__.py +++ b/src/deriver/__init__.py @@ -1,3 +1,5 @@ from .enqueue import enqueue -__all__ = ["enqueue"] +__all__ = [ + "enqueue", +] diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 4420aaa5..8b90e687 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -8,13 +8,14 @@ from rich.console import Console from src.config import settings from src.dependencies import tracked_db -from src.deriver import deriver +from src.deriver.deriver import process_representation_tasks_batch from src.utils import summarizer from src.utils.logging import log_performance_metrics from src.webhooks import webhook_delivery from .queue_payload import ( RepresentationPayload, + RepresentationPayloads, SummaryPayload, WebhookPayload, ) @@ -25,25 +26,32 @@ logging.getLogger("sqlalchemy.engine.Engine").disabled = True console = Console(markup=True) -async def process_item(task_type: str, payload: dict[str, Any]) -> None: - """Validate an incoming queue payload and dispatch it to the appropriate handler. +async def process_items(task_type: str, queue_payloads: list[dict[str, Any]]) -> None: + """Validate incoming queue payloads and dispatch to the appropriate handler. This function centralizes payload validation using a simple mapping from - task type to Pydantic model. After validation, it routes the request to + task type to Pydantic model. After validation, routes the request to the correct processor without repeating type checks elsewhere. """ - logger.debug("process_item received payload for task type %s", task_type) + logger.debug( + "process_item received %s payloads for task type %s", + len(queue_payloads), + task_type, + ) if task_type == "webhook": try: - validated = WebhookPayload(**payload) + validated = WebhookPayload(**queue_payloads[0]) except ValidationError as e: logger.error( - "Invalid webhook payload received: %s. Payload: %s", str(e), payload + "Invalid webhook payload received: %s. Payload: %s", + str(e), + queue_payloads[0], ) raise ValueError(f"Invalid payload structure: {str(e)}") from e await process_webhook(validated) logger.debug("Finished processing webhook %s", validated.event_type) + elif task_type == "summary": if settings.LANGFUSE_PUBLIC_KEY: langfuse_context.update_current_trace( # type: ignore @@ -52,13 +60,16 @@ async def process_item(task_type: str, payload: dict[str, Any]) -> None: } ) try: - validated = SummaryPayload(**payload) + validated = SummaryPayload(**queue_payloads[0]) except ValidationError as e: logger.error( - "Invalid summary payload received: %s. Payload: %s", str(e), payload + "Invalid summary payload received: %s. Payload: %s", + str(e), + queue_payloads[0], ) raise ValueError(f"Invalid payload structure: {str(e)}") from e await process_summary_task(validated) + elif task_type == "representation": if settings.LANGFUSE_PUBLIC_KEY: langfuse_context.update_current_trace( @@ -66,17 +77,22 @@ async def process_item(task_type: str, payload: dict[str, Any]) -> None: "critical_analysis_model": settings.DERIVER.MODEL, } ) - try: - validated = RepresentationPayload(**payload) + validated_payloads = RepresentationPayloads( + payloads=[ + RepresentationPayload(**payload) for payload in queue_payloads + ] + ) except ValidationError as e: logger.error( - "Invalid representation payload received: %s. Payload: %s", + "Invalid representation payloads received: %s. Payloads: %s", str(e), - payload, + queue_payloads, ) raise ValueError(f"Invalid payload structure: {str(e)}") from e - await deriver.process_representation_task(validated) + + await process_representation_tasks_batch(validated_payloads.payloads) + else: raise ValueError(f"Invalid task type: {task_type}") diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 8c35da21..684e0ff9 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -102,6 +102,16 @@ async def peer_card_call( ) +@sentry_sdk.trace +async def process_representation_tasks_batch( + payloads: list[RepresentationPayload], # pyright: ignore[reportUnusedParameter] +) -> None: + """ + Process a batch of representation tasks. + """ + pass + + @conditional_observe @sentry_sdk.trace async def process_representation_task( diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 802963d2..cf902b53 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -12,13 +12,12 @@ from sqlalchemy import delete, select, update from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql import func -from src import exceptions from src.config import settings from src.models import QueueItem -from .. import models +from .. import exceptions, models from ..dependencies import tracked_db -from .consumer import process_item +from .consumer import process_items logger = getLogger(__name__) @@ -233,59 +232,76 @@ class QueueManager: async def process_work_unit(self, work_unit_key: str): """Process all messages for a specific work unit by routing to the correct handler.""" logger.debug(f"Starting to process work unit {work_unit_key}") - async with ( - self.semaphore - ): # Hold the semaphore for the entire work unit duration + async with self.semaphore: message_count = 0 try: while not self.shutdown_event.is_set(): - message = await self.get_next_message(work_unit_key) - if not message: + candidate_messages = await self.get_message_batch( + work_unit_key, limit=20 + ) + if not candidate_messages: logger.debug(f"No more messages for work unit {work_unit_key}") break - message_count += 1 + next_message = candidate_messages[0] + task_type = next_message.task_type + + messages_to_process: list[QueueItem] = [] + + if task_type != "representation": + messages_to_process.append(next_message) + else: + # It's a representation task, build a batch + token_count = 0 + max_tokens = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + + for msg in candidate_messages: + if msg.task_type == "representation": + msg_tokens = msg.token_count or 0 + if ( + not messages_to_process + or token_count + msg_tokens <= max_tokens + ): + messages_to_process.append(msg) + token_count += msg_tokens + else: + break + else: + break + + if not messages_to_process: + logger.warning( + "No messages to process, breaking loop for work unit %s to prevent infinite loop.", + work_unit_key, + ) + break + + # Process the batch/single item try: - logger.info( - f"Processing item for task type {message.task_type} with id {message.id} from work unit {work_unit_key}" - ) - await process_item(message.task_type, message.payload) - logger.debug( - f"Successfully processed queue item for task type {message.task_type} with id {message.id}" - ) + raw_payloads = [msg.payload for msg in messages_to_process] + await process_items(task_type, raw_payloads) except exceptions.LLMError as e: logger.error( - f"LLM returned bad JSON for message {message}, re-queueing", + f"LLM returned bad JSON for messages in work unit {work_unit_key}, re-queueing", ) if settings.SENTRY.ENABLED: sentry_sdk.capture_exception(e) - continue + continue # Don't mark as processed, allow re-queue except Exception as e: logger.error( - f"Error processing queue item for task type {message.task_type} with id {message.id}: {str(e)}", + f"Error processing tasks for work unit {work_unit_key}: {e}", exc_info=True, ) if settings.SENTRY.ENABLED: sentry_sdk.capture_exception(e) - # Prevent malformed messages from stalling queue indefinitely - async with tracked_db("process_message") as db: - await db.execute( - update(models.QueueItem) - .where(models.QueueItem.id == message.id) - .values(processed=True) - ) - - await db.execute( - update(models.ActiveQueueSession) - .where( - models.ActiveQueueSession.work_unit_key == work_unit_key - ) - .values(last_updated=func.now()) - ) - - await db.commit() + # Mark messages as processed (only for non-LLM errors) + await self.mark_messages_as_processed( + messages_to_process, work_unit_key + ) + message_count += len(messages_to_process) + # Check for shutdown after processing each batch if self.shutdown_event.is_set(): logger.debug( "Shutdown requested, stopping processing for work unit %s", @@ -293,9 +309,6 @@ class QueueManager: ) break - logger.debug( - f"Completed processing work unit {work_unit_key}, processed {message_count} messages" - ) finally: # Remove work unit from active_queue_sessions when done logger.debug(f"Removing work unit {work_unit_key} from active sessions") @@ -338,22 +351,43 @@ class QueueManager: self.untrack_work_unit(work_unit_key) @sentry_sdk.trace - async def get_next_message(self, work_unit_key: str) -> QueueItem | None: - """Get the next unprocessed message for a specific work unit.""" - async with tracked_db("get_next_message") as db: + async def get_message_batch( + self, work_unit_key: str, limit: int + ) -> list[QueueItem]: + """Get a batch of unprocessed messages for a specific work unit ordered by id.""" + async with tracked_db("get_message_batch") as db: query = ( select(models.QueueItem) .where(models.QueueItem.work_unit_key == work_unit_key) .where(~models.QueueItem.processed) .order_by(models.QueueItem.id) - .limit(1) + .limit(limit) ) result = await db.execute(query) - message = result.scalar_one_or_none() + messages = result.scalars().all() # Important: commit to avoid tracked_db's rollback expiring the instance # We rely on expire_on_commit=False to keep attributes accessible post-close await db.commit() - return message + return list(messages) + + async def mark_messages_as_processed( + self, messages: list[QueueItem], work_unit_key: str + ): + if not messages: + return + async with tracked_db("process_message_batch") as db: + message_ids = [msg.id for msg in messages] + await db.execute( + update(models.QueueItem) + .where(models.QueueItem.id.in_(message_ids)) + .values(processed=True) + ) + await db.execute( + update(models.ActiveQueueSession) + .where(models.ActiveQueueSession.work_unit_key == work_unit_key) + .values(last_updated=func.now()) + ) + await db.commit() async def _cleanup_work_unit(self, work_unit_key: str) -> bool: async with tracked_db("cleanup_work_unit") as db: diff --git a/src/deriver/queue_payload.py b/src/deriver/queue_payload.py index c4f28245..b787457d 100644 --- a/src/deriver/queue_payload.py +++ b/src/deriver/queue_payload.py @@ -23,6 +23,12 @@ class RepresentationPayload(BasePayload): created_at: datetime +class RepresentationPayloads(BasePayload): + """Payload for a batch of representation tasks.""" + + payloads: list[RepresentationPayload] + + class SummaryPayload(BasePayload): """Payload for summary tasks.""" diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index e38571c0..d8d89bae 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -150,13 +150,15 @@ class TestQueueProcessing: first, second = ordered[0], ordered[1] qm = QueueManager() - nxt = await qm.get_next_message(first.work_unit_key) + batch = await qm.get_message_batch(first.work_unit_key, limit=1) + nxt = batch[0] if batch else None assert nxt is not None and nxt.id == first.id # Mark first processed, next should be the second first.processed = True await db_session.commit() - nxt2 = await qm.get_next_message(first.work_unit_key) + batch2 = await qm.get_message_batch(first.work_unit_key, limit=1) + nxt2 = batch2[0] if batch2 else None assert nxt2 is not None and nxt2.id == second.id @pytest.mark.asyncio