refactor: queue manager

This commit is contained in:
Rajat Ahuja 2025-09-05 12:18:03 -04:00
parent d09c488ed1
commit e6c580b660
7 changed files with 140 additions and 62 deletions

View File

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

View File

@ -1,3 +1,5 @@
from .enqueue import enqueue
__all__ = ["enqueue"]
__all__ = [
"enqueue",
]

View File

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

View File

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

View File

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

View File

@ -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."""

View File

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