refactor: queue manager
This commit is contained in:
parent
d09c488ed1
commit
e6c580b660
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
from .enqueue import enqueue
|
||||
|
||||
__all__ = ["enqueue"]
|
||||
__all__ = [
|
||||
"enqueue",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue