From d09c488ed1f417a0842c6e4e2bd8bd69775c9c09 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Thu, 28 Aug 2025 17:14:46 -0400 Subject: [PATCH 01/18] refactor: add token count to queue item table --- ...362d_add_token_count_to_queueitem_table.py | 113 ++++++++++++++++++ src/models.py | 1 + 2 files changed, 114 insertions(+) create mode 100644 migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py diff --git a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py new file mode 100644 index 00000000..d174917c --- /dev/null +++ b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py @@ -0,0 +1,113 @@ +"""Add token_count to QueueItem table + +Revision ID: 394e7c39362d +Revises: 88b0fb10906f +Create Date: 2025-08-27 10:49:26.591473 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +from migrations.utils import column_exists +from src.config import settings + +# revision identifiers, used by Alembic. +revision: str = "394e7c39362d" +down_revision: str | None = "88b0fb10906f" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +schema = settings.DB.SCHEMA + + +def upgrade() -> None: + op.add_column( + "queue", + sa.Column("token_count", sa.Integer(), nullable=False, server_default="0"), + schema=schema, + ) + + # ### Data Migration: Backfill token_count using raw SQL ### + bind = op.get_bind() + schema_name = settings.DB.SCHEMA + + BATCH_SIZE = 500 # Process 500 items at a time + + while True: + # Fetch a batch of items that need backfilling using raw SQL + items_result = bind.execute( + sa.text( + f""" + SELECT id, payload FROM {schema_name}.queue + WHERE processed = false + AND token_count = 0 + AND task_type IN ('representation', 'summary') + ORDER BY id + LIMIT :batch_size + """ + ), + {"batch_size": BATCH_SIZE}, + ) + items_to_backfill = items_result.fetchall() + + if not items_to_backfill: + break # No more items to process + + # Assuming message_id is always present due to application-level validation + message_id_map = { + item.id: item.payload["message_id"] for item in items_to_backfill + } + + # Get token counts for the message IDs + token_counts_result = bind.execute( + sa.text( + f""" + SELECT id, token_count FROM {schema_name}.messages + WHERE id = ANY(:ids) + """ + ), + {"ids": list(message_id_map.values())}, + ) + token_map = {msg_id: count for msg_id, count in token_counts_result.fetchall()} + + # Prepare parameters for the bulk update, skipping any queue items whose + # message_id was not found in the messages table. + update_params = [ + (qid, token_map.get(mid)) + for qid, mid in message_id_map.items() + if token_map.get(mid) is not None + ] + + if update_params: + queue_ids = [p[0] for p in update_params] + token_counts = [p[1] for p in update_params] + + # Perform a single bulk update using the UNNEST pattern + bind.execute( + sa.text( + f""" + UPDATE {schema_name}.queue q + SET token_count = v.token_count + FROM ( + SELECT + UNNEST(:queue_ids) AS queue_id, + UNNEST(:token_counts) AS token_count + ) AS v + WHERE q.id = v.queue_id + """ + ), + {"queue_ids": queue_ids, "token_counts": token_counts}, + ) + + # If we fetched fewer items than the batch size, we are on the last batch + if len(items_to_backfill) < BATCH_SIZE: + break + + +def downgrade() -> None: + inspector = sa.inspect(op.get_context().connection) + if column_exists("queue", "token_count", inspector): + op.drop_column("queue", "token_count", schema=schema) diff --git a/src/models.py b/src/models.py index defe2110..9b63d632 100644 --- a/src/models.py +++ b/src/models.py @@ -376,6 +376,7 @@ class QueueItem(Base): task_type: Mapped[TaskType] = mapped_column(TEXT, nullable=False) payload: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) processed: Mapped[bool] = mapped_column(Boolean, default=False) + token_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False) def __repr__(self) -> str: return f"QueueItem(id={self.id}, session_id={self.session_id}, work_unit_key={self.work_unit_key}, task_type={self.task_type}, payload={self.payload}, processed={self.processed})" From e6c580b660675ad11431b67f3a8773830c03f0d1 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 12:18:03 -0400 Subject: [PATCH 02/18] 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 From 4eb6830236498c6e7057360ab3fa7fc04a0366cb Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 12:46:52 -0400 Subject: [PATCH 03/18] feat: batch representation task processing --- src/deriver/deriver.py | 156 +++++++++++++++++++++++------------------ src/deriver/prompts.py | 14 ++-- tests/test_llm_mock.py | 2 +- 3 files changed, 98 insertions(+), 74 deletions(-) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 684e0ff9..39304e53 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -68,7 +68,7 @@ async def critical_analysis_call( message_created_at: datetime.datetime, working_representation: str | None, history: str, - new_turn: str, + new_turns: list[str], ): return critical_analysis_prompt( peer_id=peer_id, @@ -76,7 +76,7 @@ async def critical_analysis_call( message_created_at=message_created_at, working_representation=working_representation, history=history, - new_turn=new_turn, + new_turns=new_turns, ) @@ -104,36 +104,38 @@ async def peer_card_call( @sentry_sdk.trace async def process_representation_tasks_batch( - payloads: list[RepresentationPayload], # pyright: ignore[reportUnusedParameter] + payloads: list[RepresentationPayload], ) -> None: """ - Process a batch of representation tasks. + Process a batch of representation tasks by extracting insights and updating working representations. """ - pass + if not payloads: + return + payloads.sort(key=lambda x: x.created_at) + + latest_payload = payloads[-1] + earliest_payload = payloads[0] -@conditional_observe -@sentry_sdk.trace -async def process_representation_task( - payload: RepresentationPayload, -) -> None: - """ - Process a representation task by extracting insights and updating working representations. - """ # Start overall timing overall_start = time.perf_counter() - logger.debug("Starting insight extraction for user message: %s", payload.message_id) + logger.debug( + "Starting insight extraction for message batch starting with: %s", + earliest_payload.message_id, + ) # Use get_session_context_formatted with configurable token limit async with tracked_db("deriver.get_session_context") as db: - formatted_history = await summarizer.get_session_context_formatted( - db, - payload.workspace_name, - payload.session_name, - token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, - cutoff=payload.message_id, - include_summary=True, + formatted_history = ( + await summarizer.get_session_context_formatted( # NEED TO FIX? + db, + latest_payload.workspace_name, + latest_payload.session_name, + token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, + cutoff=latest_payload.message_id, + include_summary=True, + ) ) # instantiate embedding store from collection @@ -142,9 +144,9 @@ async def process_representation_task( # being observed by the target. collection_name = ( crud.construct_collection_name( - observer=payload.target_name, observed=payload.sender_name + observer=latest_payload.target_name, observed=latest_payload.sender_name ) - if payload.sender_name != payload.target_name + if latest_payload.sender_name != latest_payload.target_name else GLOBAL_REPRESENTATION_COLLECTION_NAME ) @@ -152,21 +154,21 @@ async def process_representation_task( async with tracked_db("deriver.get_or_create_collection") as db: collection = await crud.get_or_create_collection( db, - payload.workspace_name, + latest_payload.workspace_name, collection_name, - payload.sender_name, + latest_payload.sender_name, ) collection_name_loaded = collection.name # Use the embedding store directly embedding_store = EmbeddingStore( - workspace_name=payload.workspace_name, - peer_name=payload.sender_name, + workspace_name=latest_payload.workspace_name, + peer_name=latest_payload.sender_name, collection_name=collection_name_loaded, ) # Create reasoner instance - reasoner = CertaintyReasoner(embedding_store=embedding_store, ctx=payload) + reasoner = CertaintyReasoner(embedding_store=embedding_store, ctx=payloads) # Check for existing working representation first, fall back to global search async with tracked_db("deriver.get_working_representation_data") as db: @@ -174,10 +176,10 @@ async def process_representation_task( dict[str, Any] | str | None ) = await crud.get_working_representation_data( db, - payload.workspace_name, - payload.target_name, - payload.sender_name, - payload.session_name, + latest_payload.workspace_name, + latest_payload.target_name, + latest_payload.sender_name, + latest_payload.session_name, ) # Time context preparation @@ -210,8 +212,14 @@ async def process_representation_task( ) else: # No existing working representation, use global search + # For the first turn of a batch, we need some query text to get relevant observations. + # We'll use the content of the first message in the batch. + query_text = [payload.content for payload in payloads] + query_text = "\n".join( + query_text + ) # we probably want to think about how to handle this better working_representation = await embedding_store.get_relevant_observations( - query=payload.content, + query=query_text, conversation_context=formatted_history, for_reasoning=True, ) @@ -222,7 +230,7 @@ async def process_representation_task( logger.info("No working representation found, using global semantic search") context_prep_duration = (time.perf_counter() - context_prep_start) * 1000 accumulate_metric( - f"deriver_representation_{payload.message_id}_{payload.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "context_preparation", context_prep_duration, "ms", @@ -235,10 +243,13 @@ async def process_representation_task( async with tracked_db("deriver.get_peer_card") as db: speaker_peer_card: list[str] | None = await crud.get_peer_card( - db, payload.workspace_name, payload.sender_name, payload.target_name + db, + latest_payload.workspace_name, + latest_payload.sender_name, + latest_payload.target_name, ) if speaker_peer_card is None: - logger.warning("No peer card found for %s", payload.sender_name) + logger.warning("No peer card found for %s", latest_payload.sender_name) else: logger.info("Using peer card: %s", speaker_peer_card) @@ -247,6 +258,7 @@ async def process_representation_task( working_representation, formatted_history, speaker_peer_card, + payloads, ) logger.debug("REASONING COMPLETION: Unified reasoning completed across all levels.") @@ -258,12 +270,11 @@ async def process_representation_task( log_observations_tree(final_obs_dict) # Always save working representation to peer for dialectic access - await save_working_representation_to_peer(payload, final_observations) - + await save_working_representation_to_peer(latest_payload, final_observations) # Calculate and log overall timing overall_duration = (time.perf_counter() - overall_start) * 1000 accumulate_metric( - f"deriver_representation_{payload.message_id}_{payload.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "total_processing_time", overall_duration, "ms", @@ -272,13 +283,13 @@ async def process_representation_task( total_observations = sum(len(obs_list) for obs_list in final_obs_dict.values()) accumulate_metric( - f"deriver_representation_{payload.message_id}_{payload.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "final_observation_count", total_observations, "", ) log_performance_metrics( - f"deriver_representation_{payload.message_id}_{payload.target_name}" + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}" ) if settings.LANGFUSE_PUBLIC_KEY: @@ -287,14 +298,21 @@ async def process_representation_task( ) +# The old function now just calls the batch processor with a single payload +async def process_representation_task( + payload: RepresentationPayload, +) -> None: + await process_representation_tasks_batch([payload]) + + class CertaintyReasoner: """Certainty reasoner for analyzing and deriving insights.""" embedding_store: EmbeddingStore - ctx: RepresentationPayload + ctx: list[RepresentationPayload] def __init__( - self, embedding_store: EmbeddingStore, ctx: RepresentationPayload + self, embedding_store: EmbeddingStore, ctx: list[RepresentationPayload] ) -> None: self.embedding_store = embedding_store self.ctx = ctx @@ -310,47 +328,49 @@ class CertaintyReasoner: """ Critically analyzes and revises understanding, returning structured observations. """ + # For logging, we can just show the content of the last message + latest_payload = self.ctx[-1] if settings.LANGFUSE_PUBLIC_KEY: langfuse_context.update_current_observation( input=format_reasoning_inputs_as_markdown( working_representation, history, - self.ctx.content, - self.ctx.created_at, + latest_payload.content, + latest_payload.created_at, ) ) - formatted_new_turn = format_new_turn_with_timestamp( - self.ctx.content, - self.ctx.created_at, - self.ctx.sender_name, - ) + new_turns = [ + format_new_turn_with_timestamp(p.content, p.created_at, p.sender_name) + for p in self.ctx + ] + formatted_working_representation = format_context_for_prompt( working_representation ) logger.debug( - "CRITICAL ANALYSIS: message_created_at='%s', formatted_new_turn='%s'", - self.ctx.created_at, - formatted_new_turn, + "CRITICAL ANALYSIS: message_created_at='%s', new_turns_count=%s", + latest_payload.created_at, + len(new_turns), ) try: response_obj = await critical_analysis_call( - peer_id=self.ctx.sender_name, + peer_id=latest_payload.sender_name, peer_card=speaker_peer_card, - message_created_at=self.ctx.created_at, + message_created_at=latest_payload.created_at, working_representation=formatted_working_representation, history=history, - new_turn=formatted_new_turn, + new_turns=new_turns, ) except Exception as e: raise exceptions.LLMError( speaker_peer_card=speaker_peer_card, working_representation=formatted_working_representation, history=history, - new_turn=formatted_new_turn, + new_turns=new_turns, ) from e # If response is a string, try to parse as JSON @@ -420,6 +440,7 @@ class CertaintyReasoner: Single-pass reasoning function that critically analyzes and derives insights. Performs one analysis pass and returns the final observations. """ + latest_payload = self.ctx[-1] analysis_start = time.perf_counter() # Perform critical analysis to get observation lists @@ -434,7 +455,7 @@ class CertaintyReasoner: analysis_duration_ms = (time.perf_counter() - analysis_start) * 1000 accumulate_metric( - f"deriver_representation_{self.ctx.message_id}_{self.ctx.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "critical_analysis_duration", analysis_duration_ms, "ms", @@ -445,13 +466,13 @@ class CertaintyReasoner: new_observations_by_level: dict[ str, list[str] ] = await self._save_new_observations( - working_representation, reasoning_response + working_representation, reasoning_response, latest_payload ) save_observations_duration = ( time.perf_counter() - save_observations_start ) * 1000 accumulate_metric( - f"deriver_representation_{self.ctx.message_id}_{self.ctx.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "save_new_observations", save_observations_duration, "ms", @@ -470,7 +491,7 @@ class CertaintyReasoner: time.perf_counter() - update_peer_card_start ) * 1000 accumulate_metric( - f"deriver_representation_{self.ctx.message_id}_{self.ctx.target_name}", + f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "update_peer_card", update_peer_card_duration, "ms", @@ -484,6 +505,7 @@ class CertaintyReasoner: self, original_working_representation: ReasoningResponse, revised_observations: ReasoningResponse, + latest_payload: RepresentationPayload, ) -> dict[str, list[str]]: """Save only the observations that are new compared to the original context.""" # Use the utility function to find new observations @@ -531,9 +553,9 @@ class CertaintyReasoner: if all_unified_observations: await self.embedding_store.save_unified_observations( all_unified_observations, - self.ctx.message_id, - self.ctx.session_name, - self.ctx.created_at, + latest_payload.message_id, + latest_payload.session_name, + latest_payload.created_at, ) else: logger.debug("No new observations to save") @@ -567,9 +589,9 @@ class CertaintyReasoner: async with tracked_db("deriver.update_peer_card") as db: await crud.set_peer_card( db, - self.ctx.workspace_name, - self.ctx.sender_name, - self.ctx.target_name, + self.ctx[0].workspace_name, + self.ctx[0].sender_name, + self.ctx[0].target_name, new_peer_card, ) except Exception as e: diff --git a/src/deriver/prompts.py b/src/deriver/prompts.py index 0746895c..27352340 100644 --- a/src/deriver/prompts.py +++ b/src/deriver/prompts.py @@ -18,7 +18,7 @@ def critical_analysis_prompt( message_created_at: datetime.datetime, working_representation: str | None, history: str, - new_turn: str, + new_turns: list[str], ) -> str: """ Generate the critical analysis prompt for the deriver. @@ -29,7 +29,7 @@ def critical_analysis_prompt( message_created_at (datetime.datetime): Timestamp of the message. working_representation (str | None): Current user understanding context. history (str): Recent conversation history. - new_turn (str): New conversation turn to analyze. + new_turns (list[str]): New conversation turns to analyze. Returns: Formatted prompt string for critical analysis @@ -58,6 +58,8 @@ The current user understanding: else "" ) + new_turns_section = "\n".join(new_turns) + return c( f""" You are an agent who critically analyzes user messages through rigorous logical reasoning to produce only conclusions about the user that are CERTAIN. @@ -94,10 +96,10 @@ Recent conversation history for context: {history} -New conversation turn to analyze: - -{new_turn} - +New conversation turns to analyze: + +{new_turns_section} + """ ) diff --git a/tests/test_llm_mock.py b/tests/test_llm_mock.py index 55fc55dd..306b8c31 100644 --- a/tests/test_llm_mock.py +++ b/tests/test_llm_mock.py @@ -20,7 +20,7 @@ async def test_generic_honcho_llm_call_mock(): message_created_at=datetime(2023, 1, 1, 0, 0, 0, tzinfo=timezone.utc), working_representation="test working representation", history="test history", - new_turn="test new turn", + new_turns=["test new turn"], ) # Verify that we get a mock result, not an actual LLM call From 4328fb3d2abbec91658f06c79d28caccf870b78c Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 14:33:42 -0400 Subject: [PATCH 04/18] fix: actually pass in token count from msg to queue item --- src/deriver/enqueue.py | 3 ++- src/deriver/queue_manager.py | 20 +++++++++----------- src/routers/messages.py | 2 ++ 3 files changed, 13 insertions(+), 12 deletions(-) diff --git a/src/deriver/enqueue.py b/src/deriver/enqueue.py index 00028cda..5974fc9e 100644 --- a/src/deriver/enqueue.py +++ b/src/deriver/enqueue.py @@ -102,7 +102,6 @@ async def handle_session( message_seq_map=message_seq_map, ) ) - return queue_records @@ -166,6 +165,7 @@ def create_representation_record( "payload": processed_payload, "session_id": session_id, "task_type": "representation", + "token_count": message.get("token_count"), } @@ -198,6 +198,7 @@ def create_summary_record( "payload": processed_payload, "session_id": session_id, "task_type": "summary", + "token_count": message.get("token_count"), } diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index cf902b53..6f2f683c 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -237,7 +237,8 @@ class QueueManager: try: while not self.shutdown_event.is_set(): candidate_messages = await self.get_message_batch( - work_unit_key, limit=20 + work_unit_key, + limit=10, # hard limit of 10 messages per batch ) if not candidate_messages: logger.debug(f"No more messages for work unit {work_unit_key}") @@ -256,16 +257,13 @@ class QueueManager: 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 + 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 diff --git a/src/routers/messages.py b/src/routers/messages.py index 3c59dbb3..27d89a14 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -66,6 +66,7 @@ async def create_messages_for_session( "content": message.content, "peer_name": message.peer_name, "created_at": message.created_at, + "token_count": message.token_count, } for message in created_messages ] @@ -127,6 +128,7 @@ async def create_messages_with_file( "content": message.content, "peer_name": message.peer_name, "created_at": message.created_at, + "token_count": message.token_count, } for message in created_messages ] From faefb9ac867a0dc1fc27a7845492fe292c8d0600 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 16:12:35 -0400 Subject: [PATCH 05/18] test: queue manager --- tests/deriver/test_queue_processing.py | 219 +++++++++++++++++++++++++ 1 file changed, 219 insertions(+) diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index d8d89bae..788c36a8 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -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) From 6ae0d9f4344e0ac9b78d3b5cc55e6810a3f455f2 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 5 Sep 2025 16:30:55 -0400 Subject: [PATCH 06/18] chore: CR comments --- src/deriver/consumer.py | 2 +- src/deriver/deriver.py | 5 +++-- tests/deriver/test_queue_processing.py | 19 ++++++++++++++----- 3 files changed, 18 insertions(+), 8 deletions(-) diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 8b90e687..5e3de9ea 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -34,7 +34,7 @@ async def process_items(task_type: str, queue_payloads: list[dict[str, Any]]) -> the correct processor without repeating type checks elsewhere. """ logger.debug( - "process_item received %s payloads for task type %s", + "process_items received %s payloads for task type %s", len(queue_payloads), task_type, ) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 39304e53..fe987ef5 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -503,8 +503,9 @@ class CertaintyReasoner: @sentry_sdk.trace async def _save_new_observations( self, - original_working_representation: ReasoningResponse, - revised_observations: ReasoningResponse, + original_working_representation: ReasoningResponse + | ReasoningResponseWithThinking, + revised_observations: ReasoningResponse | ReasoningResponseWithThinking, latest_payload: RepresentationPayload, ) -> dict[str, list[str]]: """Save only the observations that are new compared to the original context.""" diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 788c36a8..746c900d 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -5,6 +5,7 @@ import pytest from sqlalchemy.ext.asyncio import AsyncSession from src import models +from src.config import settings from src.deriver.queue_manager import QueueManager @@ -251,7 +252,7 @@ class TestQueueProcessing: peer = peers[0] # Create messages with token counts that exceed batch limit - # Total: 2000 + 2000 + 3000 = 7000 tokens (> 4096 limit defined in settings) + limit = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS messages = [ models.Message( id=i, @@ -261,7 +262,7 @@ class TestQueueProcessing: content=f"Test message {i}", token_count=token_count, ) - for i, token_count in enumerate([2000, 2000, 3000]) + for i, token_count in enumerate([limit // 2, limit // 2, limit // 2]) ] # Create queue items with token counts @@ -320,8 +321,9 @@ class TestQueueProcessing: # 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 + assert processed_batches[0]["payload_count"] == 2 + assert processed_batches[1]["payload_count"] == 1 + assert all(b["task_type"] == "representation" for b in processed_batches) @pytest.mark.asyncio async def test_hard_batch_size_limit( @@ -342,7 +344,7 @@ class TestQueueProcessing: workspace_name=session.workspace_name, peer_name=peer.name, content=f"Test message {i}", - token_count=100, # Small tokens to avoid token-based batching + token_count=5, # Small tokens to avoid token-based batching ) for i in range(15) ] @@ -378,6 +380,13 @@ class TestQueueProcessing: batch = await qm.get_message_batch(queue_items[0].work_unit_key, limit=10) assert len(batch) == 10 # Should respect hard limit + # Mark the first batch as processed + await qm.mark_messages_as_processed(batch, queue_items[0].work_unit_key) + + # Get the next batch - should return remaining 5 items + next_batch = await qm.get_message_batch(queue_items[0].work_unit_key, limit=10) + assert len(next_batch) == 5 # Should return remaining items + @pytest.mark.asyncio async def test_single_message_processing( self, From 530191484791925891e3c89b16b11013046aa3e7 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Fri, 19 Sep 2025 12:24:01 -0400 Subject: [PATCH 07/18] fix: CR comments --- src/config.py | 8 ++++++++ src/deriver/consumer.py | 4 ++++ src/deriver/deriver.py | 8 +++----- src/deriver/queue_manager.py | 8 ++++---- tests/deriver/test_queue_processing.py | 28 ++++++++++++++++++++++++++ 5 files changed, 47 insertions(+), 9 deletions(-) diff --git a/src/config.py b/src/config.py index 9272206e..ad161c98 100644 --- a/src/config.py +++ b/src/config.py @@ -220,6 +220,14 @@ class DeriverSettings(HonchoSettings): ), ] = 4096 + @model_validator(mode="after") + def validate_batch_tokens_vs_context_limit(self): + if self.REPRESENTATION_BATCH_MAX_TOKENS > self.CONTEXT_TOKEN_LIMIT: + raise ValueError( + f"REPRESENTATION_BATCH_MAX_TOKENS ({self.REPRESENTATION_BATCH_MAX_TOKENS}) cannot exceed CONTEXT_TOKEN_LIMIT ({self.CONTEXT_TOKEN_LIMIT})" + ) + return self + class DialecticSettings(HonchoSettings): model_config = SettingsConfigDict(env_prefix="DIALECTIC_", extra="ignore") # pyright: ignore diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 5e3de9ea..924eb033 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -33,6 +33,10 @@ async def process_items(task_type: str, queue_payloads: list[dict[str, Any]]) -> task type to Pydantic model. After validation, routes the request to the correct processor without repeating type checks elsewhere. """ + if not queue_payloads: + logger.debug("process_items received no payloads for task type %s", task_type) + return + logger.debug( "process_items received %s payloads for task type %s", len(queue_payloads), diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index fe987ef5..0531e15b 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -112,7 +112,7 @@ async def process_representation_tasks_batch( if not payloads: return - payloads.sort(key=lambda x: x.created_at) + payloads.sort(key=lambda x: x.message_id) latest_payload = payloads[-1] earliest_payload = payloads[0] @@ -212,12 +212,10 @@ async def process_representation_tasks_batch( ) else: # No existing working representation, use global search - # For the first turn of a batch, we need some query text to get relevant observations. - # We'll use the content of the first message in the batch. query_text = [payload.content for payload in payloads] query_text = "\n".join( query_text - ) # we probably want to think about how to handle this better + ) # TODO: consider a smarter strategy than concatenation working_representation = await embedding_store.get_relevant_observations( query=query_text, conversation_context=formatted_history, @@ -286,7 +284,7 @@ async def process_representation_tasks_batch( f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}", "final_observation_count", total_observations, - "", + "count", ) log_performance_metrics( f"deriver_representation_{latest_payload.message_id}_{latest_payload.target_name}" diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 3a57a3dc..660165e3 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 models from src.config import settings +from src.dependencies import tracked_db +from src.deriver.consumer import process_items from src.models import QueueItem -from .. import models -from ..dependencies import tracked_db -from .consumer import process_items - logger = getLogger(__name__) load_dotenv(override=True) @@ -258,6 +257,7 @@ class QueueManager: for msg in candidate_messages: msg_tokens = msg.token_count or 0 + # Always process at least one message, even if over limit if ( not messages_to_process or token_count + msg_tokens <= max_tokens diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 746c900d..8975d5f6 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -464,3 +464,31 @@ class TestQueueProcessing: 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) + + # Verify the corresponding DB records are marked as processed + from sqlalchemy import select + + # Query for the summary queue items that were processed + processed_items = ( + ( + await db_session.execute( + select(models.QueueItem) + .where(models.QueueItem.work_unit_key == work_unit_key) + .where(models.QueueItem.task_type == "summary") + .order_by(models.QueueItem.id) + ) + ) + .scalars() + .all() + ) + + # Assert we found both summary items + assert len(processed_items) == 2 + + # Assert both items are marked as processed + assert all(item.processed is True for item in processed_items) + + # Optionally verify the items have the expected token counts from the messages + expected_token_counts = [500, 600] # From the test messages + actual_token_counts = [item.token_count for item in processed_items] + assert sorted(actual_token_counts) == sorted(expected_token_counts) From f4fefd5d701ff602c34dc612a209136d094b392b Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 22 Sep 2025 12:44:37 -0400 Subject: [PATCH 08/18] fix: remove payloads from reasoner.reason call --- src/deriver/deriver.py | 1 - 1 file changed, 1 deletion(-) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 0531e15b..58a59392 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -256,7 +256,6 @@ async def process_representation_tasks_batch( working_representation, formatted_history, speaker_peer_card, - payloads, ) logger.debug("REASONING COMPLETION: Unified reasoning completed across all levels.") From 1a7bd91cddb549de4c9ea528801bbd7e14bbeb57 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 22 Sep 2025 13:13:56 -0400 Subject: [PATCH 09/18] fix: use earliest message as cutoff --- src/deriver/deriver.py | 16 +++-- tests/deriver/test_deriver_processing.py | 76 ++++++++++++++++++++++++ 2 files changed, 83 insertions(+), 9 deletions(-) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 58a59392..d51c3e00 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -127,15 +127,13 @@ async def process_representation_tasks_batch( # Use get_session_context_formatted with configurable token limit async with tracked_db("deriver.get_session_context") as db: - formatted_history = ( - await summarizer.get_session_context_formatted( # NEED TO FIX? - db, - latest_payload.workspace_name, - latest_payload.session_name, - token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, - cutoff=latest_payload.message_id, - include_summary=True, - ) + formatted_history = await summarizer.get_session_context_formatted( + db, + latest_payload.workspace_name, + latest_payload.session_name, + token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, + cutoff=earliest_payload.message_id, + include_summary=True, ) # instantiate embedding store from collection diff --git a/tests/deriver/test_deriver_processing.py b/tests/deriver/test_deriver_processing.py index deaa085a..ff8b8caf 100644 --- a/tests/deriver/test_deriver_processing.py +++ b/tests/deriver/test_deriver_processing.py @@ -1,10 +1,15 @@ import signal from collections.abc import Callable, Generator +from datetime import datetime, timedelta, timezone from typing import Any +from unittest.mock import AsyncMock import pytest from src import models +from src.deriver.deriver import process_representation_tasks_batch +from src.deriver.queue_payload import RepresentationPayload +from src.utils.shared_models import ReasoningResponseWithThinking @pytest.mark.asyncio @@ -98,3 +103,74 @@ class TestDeriverProcessing: # Verify the methods were called assert mock_embedding_store.save_unified_observations.called # type: ignore[attr-defined] + + async def test_representation_batch_uses_earliest_cutoff( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Ensure batching history cutoff uses the earliest payload in the batch.""" + captured_cutoffs: list[int] = [] + + async def fake_get_session_context_formatted(*_args: Any, **kwargs: Any) -> str: + captured_cutoffs.append(kwargs["cutoff"]) + return "formatted-history" + + # Mock only the function we need to inspect for the test assertion + monkeypatch.setattr( + "src.deriver.deriver.summarizer.get_session_context_formatted", + fake_get_session_context_formatted, + ) + + # Provide a stub working representation so embedding lookups are skipped. + monkeypatch.setattr( + "src.deriver.deriver.crud.get_working_representation_data", + AsyncMock( + return_value={ + "final_observations": { + "explicit": ["existing"], + "deductive": [], + } + } + ), + ) + + # Avoid executing the full reasoning pipeline; we only care about cutoff behavior. + monkeypatch.setattr( + "src.deriver.deriver.CertaintyReasoner.reason", + AsyncMock( + return_value=ReasoningResponseWithThinking( + thinking=None, explicit=[], deductive=[] + ) + ), + ) + + # Skip persisting results back to the database. + monkeypatch.setattr( + "src.deriver.deriver.save_working_representation_to_peer", + AsyncMock(), + ) + + # Create test payloads with different message IDs (earlier message has lower ID) + now = datetime.now(timezone.utc) + payloads: list[RepresentationPayload] = [] + for i in range(8): + message_id = 100 + i # 100, 101, 102, ..., 107 + payloads.append( + RepresentationPayload( + workspace_name="test_workspace", + session_name="test_session", + message_id=message_id, + content=f"message {message_id}", + sender_name="alice", + target_name="alice", + created_at=now + - timedelta( + minutes=7 - i + ), # Earlier messages have earlier timestamps + ) + ) + + await process_representation_tasks_batch(payloads) + + # Verify that the earliest message ID was used as the cutoff + assert captured_cutoffs == [payloads[0].message_id] From 26f975f808b436dd508c4b1990ef24dc303cb7f9 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 22 Sep 2025 13:31:33 -0400 Subject: [PATCH 10/18] test: CR comments --- tests/deriver/test_deriver_processing.py | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/tests/deriver/test_deriver_processing.py b/tests/deriver/test_deriver_processing.py index ff8b8caf..020d9429 100644 --- a/tests/deriver/test_deriver_processing.py +++ b/tests/deriver/test_deriver_processing.py @@ -134,6 +134,24 @@ class TestDeriverProcessing: ), ) + # Avoid DB access for collection and peer card + monkeypatch.setattr( + "src.deriver.deriver.crud.get_or_create_collection", + AsyncMock(return_value=type("Collection", (), {"name": "dummy"})()), + ) + monkeypatch.setattr( + "src.deriver.deriver.crud.get_peer_card", + AsyncMock(return_value=[]), + ) + # Short-circuit tracked_db context manager + from contextlib import asynccontextmanager + + @asynccontextmanager + async def _no_db(_label: str): + yield object() + + monkeypatch.setattr("src.deriver.deriver.tracked_db", _no_db) + # Avoid executing the full reasoning pipeline; we only care about cutoff behavior. monkeypatch.setattr( "src.deriver.deriver.CertaintyReasoner.reason", From b412f595b0de32d3a25c94c623c32a2cc52a44a7 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 22 Sep 2025 13:59:47 -0400 Subject: [PATCH 11/18] fix: migration per CR comments --- ...362d_add_token_count_to_queueitem_table.py | 97 +++++++++++-------- 1 file changed, 59 insertions(+), 38 deletions(-) diff --git a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py index d174917c..5886aebb 100644 --- a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py +++ b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py @@ -6,6 +6,7 @@ Create Date: 2025-08-27 10:49:26.591473 """ +import json from collections.abc import Sequence import sqlalchemy as sa @@ -38,39 +39,67 @@ def upgrade() -> None: while True: # Fetch a batch of items that need backfilling using raw SQL - items_result = bind.execute( - sa.text( - f""" - SELECT id, payload FROM {schema_name}.queue - WHERE processed = false - AND token_count = 0 - AND task_type IN ('representation', 'summary') - ORDER BY id - LIMIT :batch_size - """ - ), - {"batch_size": BATCH_SIZE}, + items_stmt = sa.text( + f""" + SELECT id, payload FROM {schema_name}.queue + WHERE processed = false + AND token_count = 0 + AND task_type IN ('representation', 'summary') + ORDER BY id + LIMIT :batch_size + """ + ).columns(id=sa.Integer, payload=sa.JSON) + + items_to_backfill = ( + bind.execute( + items_stmt, + {"batch_size": BATCH_SIZE}, + ) + .mappings() + .all() ) - items_to_backfill = items_result.fetchall() if not items_to_backfill: break # No more items to process # Assuming message_id is always present due to application-level validation + def _extract_message_id(payload: object) -> int | None: + if isinstance(payload, dict): + return payload.get("message_id") # type: ignore[return-value] + if payload is None: + return None + try: + parsed = json.loads(payload) + except (TypeError, json.JSONDecodeError): + return None + if isinstance(parsed, dict): + return parsed.get("message_id") # type: ignore[return-value] + return None + message_id_map = { - item.id: item.payload["message_id"] for item in items_to_backfill + row["id"]: _extract_message_id(row["payload"]) for row in items_to_backfill + } + # Drop rows where message_id couldn't be extracted + message_id_map = { + queue_id: message_id + for queue_id, message_id in message_id_map.items() + if message_id is not None } # Get token counts for the message IDs - token_counts_result = bind.execute( - sa.text( - f""" - SELECT id, token_count FROM {schema_name}.messages - WHERE id = ANY(:ids) - """ - ), - {"ids": list(message_id_map.values())}, + token_ids = list(message_id_map.values()) + if not token_ids: + continue + + placeholders = ", ".join(f":id_{idx}" for idx in range(len(token_ids))) + token_counts_stmt = sa.text( + f""" + SELECT id, token_count FROM {schema_name}.messages + WHERE id IN ({placeholders}) + """ ) + bind_params = {f"id_{idx}": token_id for idx, token_id in enumerate(token_ids)} + token_counts_result = bind.execute(token_counts_stmt, bind_params) token_map = {msg_id: count for msg_id, count in token_counts_result.fetchall()} # Prepare parameters for the bulk update, skipping any queue items whose @@ -82,24 +111,16 @@ def upgrade() -> None: ] if update_params: - queue_ids = [p[0] for p in update_params] - token_counts = [p[1] for p in update_params] + update_stmt = sa.text( + f"UPDATE {schema_name}.queue SET token_count = :token_count WHERE id = :queue_id" + ) - # Perform a single bulk update using the UNNEST pattern bind.execute( - sa.text( - f""" - UPDATE {schema_name}.queue q - SET token_count = v.token_count - FROM ( - SELECT - UNNEST(:queue_ids) AS queue_id, - UNNEST(:token_counts) AS token_count - ) AS v - WHERE q.id = v.queue_id - """ - ), - {"queue_ids": queue_ids, "token_counts": token_counts}, + update_stmt, + [ + {"queue_id": queue_id, "token_count": token_count} + for queue_id, token_count in update_params + ], ) # If we fetched fewer items than the batch size, we are on the last batch From 13574e1587e6887e5708f3e0cd7c1ee57f25f032 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 22 Sep 2025 14:05:31 -0400 Subject: [PATCH 12/18] fix: nit issues --- ...362d_add_token_count_to_queueitem_table.py | 24 +++++++++---------- 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py index 5886aebb..306c772b 100644 --- a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py +++ b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py @@ -31,20 +31,20 @@ def upgrade() -> None: schema=schema, ) - # ### Data Migration: Backfill token_count using raw SQL ### bind = op.get_bind() - schema_name = settings.DB.SCHEMA BATCH_SIZE = 500 # Process 500 items at a time + last_seen_id: int = 0 while True: - # Fetch a batch of items that need backfilling using raw SQL items_stmt = sa.text( f""" - SELECT id, payload FROM {schema_name}.queue + SELECT id, payload FROM {schema}.queue WHERE processed = false AND token_count = 0 AND task_type IN ('representation', 'summary') + AND (payload->>'message_id') IS NOT NULL + AND id > :after_id ORDER BY id LIMIT :batch_size """ @@ -53,7 +53,7 @@ def upgrade() -> None: items_to_backfill = ( bind.execute( items_stmt, - {"batch_size": BATCH_SIZE}, + {"batch_size": BATCH_SIZE, "after_id": last_seen_id}, ) .mappings() .all() @@ -61,19 +61,19 @@ def upgrade() -> None: if not items_to_backfill: break # No more items to process + last_seen_id = items_to_backfill[-1]["id"] - # Assuming message_id is always present due to application-level validation def _extract_message_id(payload: object) -> int | None: if isinstance(payload, dict): - return payload.get("message_id") # type: ignore[return-value] + return payload.get("message_id") # pyright: ignore if payload is None: return None try: - parsed = json.loads(payload) + parsed = json.loads(str(payload)) except (TypeError, json.JSONDecodeError): return None if isinstance(parsed, dict): - return parsed.get("message_id") # type: ignore[return-value] + return parsed.get("message_id") # pyright: ignore return None message_id_map = { @@ -94,7 +94,7 @@ def upgrade() -> None: placeholders = ", ".join(f":id_{idx}" for idx in range(len(token_ids))) token_counts_stmt = sa.text( f""" - SELECT id, token_count FROM {schema_name}.messages + SELECT id, token_count FROM {schema}.messages WHERE id IN ({placeholders}) """ ) @@ -112,7 +112,7 @@ def upgrade() -> None: if update_params: update_stmt = sa.text( - f"UPDATE {schema_name}.queue SET token_count = :token_count WHERE id = :queue_id" + f"UPDATE {schema}.queue SET token_count = :token_count WHERE id = :queue_id" ) bind.execute( @@ -129,6 +129,6 @@ def upgrade() -> None: def downgrade() -> None: - inspector = sa.inspect(op.get_context().connection) + inspector = sa.inspect(op.get_bind()) if column_exists("queue", "token_count", inspector): op.drop_column("queue", "token_count", schema=schema) From 561e4f67436ff48451e6be0ce5a6db844a30f65c Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Tue, 23 Sep 2025 11:25:21 -0400 Subject: [PATCH 13/18] fix: config.toml and .env.template --- .env.template | 13 ++++++++++++- config.toml.example | 9 +++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/.env.template b/.env.template index abb4e5b3..9c82815d 100644 --- a/.env.template +++ b/.env.template @@ -10,6 +10,8 @@ LOG_LEVEL=INFO # SESSION_OBSERVERS_LIMIT=10 # GET_CONTEXT_MAX_TOKENS=100000 +# MAX_FILE_SIZE=5242880 +# MAX_MESSAGE_SIZE=25000 # Embedding settings # EMBED_MESSAGES=true @@ -89,6 +91,7 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # DERIVER_PEER_CARD_MAX_OUTPUT_TOKENS=2000 # DERIVER_CONTEXT_TOKEN_LIMIT=30000 # DERIVER_WORKING_REPRESENTATION_MAX_OBSERVATIONS=100 +# DERIVER_REPRESENTATION_BATCH_MAX_TOKENS=4096 # ============================================================================= # Dialectic Settings @@ -102,6 +105,7 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # DIALECTIC_SEMANTIC_SEARCH_TOP_K=10 # DIALECTIC_SEMANTIC_SEARCH_MAX_DISTANCE=0.85 # DIALECTIC_THINKING_BUDGET_TOKENS=1024 +# DIALECTIC_CONTEXT_WINDOW_SIZE=100000 # ============================================================================= # Summary Settings @@ -112,6 +116,13 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # SUMMARY_MODEL=gemini-1.5-flash-latest # SUMMARY_MAX_TOKENS_SHORT=1000 # SUMMARY_MAX_TOKENS_LONG=2000 +# SUMMARY_THINKING_BUDGET_TOKENS=512 + +# ============================================================================= +# Webhook Settings +# ============================================================================= +# WEBHOOK_SECRET= +# WEBHOOK_MAX_WORKSPACE_LIMIT=10 # ============================================================================= # Monitoring and Observability (Optional) @@ -122,4 +133,4 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # SENTRY_RELEASE=your-release-semver # SENTRY_ENVIRONMENT=development # SENTRY_TRACES_SAMPLE_RATE=0.1 -# SENTRY_PROFILES_SAMPLE_RATE=0.1 +# SENTRY_PROFILES_SAMPLE_RATE=0.1 \ No newline at end of file diff --git a/config.toml.example b/config.toml.example index 98d84f5a..982f3561 100644 --- a/config.toml.example +++ b/config.toml.example @@ -8,6 +8,8 @@ LOG_LEVEL = "INFO" SESSION_OBSERVERS_LIMIT = 10 GET_CONTEXT_MAX_TOKENS = 100000 +MAX_FILE_SIZE = 5242880 # 5MB +MAX_MESSAGE_SIZE = 25000 EMBED_MESSAGES = true MAX_EMBEDDING_TOKENS = 8192 MAX_EMBEDDING_TOKENS_PER_REQUEST = 300000 @@ -69,6 +71,7 @@ PEER_CARD_MODEL = "gpt-5-nano-2025-08-07" PEER_CARD_MAX_OUTPUT_TOKENS = 2000 CONTEXT_TOKEN_LIMIT = 30000 WORKING_REPRESENTATION_MAX_OBSERVATIONS = 100 +REPRESENTATION_BATCH_MAX_TOKENS = 4096 # Dialectic settings [dialectic] @@ -81,6 +84,7 @@ MAX_OUTPUT_TOKENS = 2500 SEMANTIC_SEARCH_TOP_K = 10 SEMANTIC_SEARCH_MAX_DISTANCE = 0.85 THINKING_BUDGET_TOKENS = 1024 +CONTEXT_WINDOW_SIZE = 100000 # Summary settings [summary] @@ -91,3 +95,8 @@ MODEL = "gemini-1.5-flash-latest" MAX_TOKENS_SHORT = 1000 MAX_TOKENS_LONG = 2000 THINKING_BUDGET_TOKENS = 512 + +# Webhook settings +[webhook] +SECRET = "" +MAX_WORKSPACE_LIMIT = 10 From 0b58472ab1621da901c23155b993935beedf5d34 Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Tue, 23 Sep 2025 17:25:11 -0400 Subject: [PATCH 14/18] chore: Coderabbit Nitpicks --- .env.template | 6 +++--- config.toml.example | 2 +- src/deriver/consumer.py | 2 +- src/deriver/deriver.py | 2 +- 4 files changed, 6 insertions(+), 6 deletions(-) diff --git a/.env.template b/.env.template index 9c82815d..9cc57de0 100644 --- a/.env.template +++ b/.env.template @@ -10,8 +10,8 @@ LOG_LEVEL=INFO # SESSION_OBSERVERS_LIMIT=10 # GET_CONTEXT_MAX_TOKENS=100000 -# MAX_FILE_SIZE=5242880 -# MAX_MESSAGE_SIZE=25000 +# MAX_FILE_SIZE=5242880 # Bytes +# MAX_MESSAGE_SIZE=25000 # Characters # Embedding settings # EMBED_MESSAGES=true @@ -133,4 +133,4 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # SENTRY_RELEASE=your-release-semver # SENTRY_ENVIRONMENT=development # SENTRY_TRACES_SAMPLE_RATE=0.1 -# SENTRY_PROFILES_SAMPLE_RATE=0.1 \ No newline at end of file +# SENTRY_PROFILES_SAMPLE_RATE=0.1 diff --git a/config.toml.example b/config.toml.example index 982f3561..3be89786 100644 --- a/config.toml.example +++ b/config.toml.example @@ -9,7 +9,7 @@ LOG_LEVEL = "INFO" SESSION_OBSERVERS_LIMIT = 10 GET_CONTEXT_MAX_TOKENS = 100000 MAX_FILE_SIZE = 5242880 # 5MB -MAX_MESSAGE_SIZE = 25000 +MAX_MESSAGE_SIZE = 25000 # Characters EMBED_MESSAGES = true MAX_EMBEDDING_TOKENS = 8192 MAX_EMBEDDING_TOKENS_PER_REQUEST = 300000 diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 924eb033..a6796752 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -33,7 +33,7 @@ async def process_items(task_type: str, queue_payloads: list[dict[str, Any]]) -> task type to Pydantic model. After validation, routes the request to the correct processor without repeating type checks elsewhere. """ - if not queue_payloads: + if not queue_payloads or not queue_payloads[0]: logger.debug("process_items received no payloads for task type %s", task_type) return diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index d51c3e00..f7352dd3 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -109,7 +109,7 @@ async def process_representation_tasks_batch( """ Process a batch of representation tasks by extracting insights and updating working representations. """ - if not payloads: + if not payloads or len(payloads) == 0: return payloads.sort(key=lambda x: x.message_id) From 1cc3d9aaaf6e56c956796ead4dd5dd7fad793ffc Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 24 Sep 2025 11:45:04 -0400 Subject: [PATCH 15/18] perf: use join on messages table instead of storing token_count in queue payload --- ...362d_add_token_count_to_queueitem_table.py | 134 ------- src/deriver/enqueue.py | 2 - src/deriver/queue_manager.py | 131 ++++--- src/models.py | 1 - tests/deriver/test_queue_processing.py | 345 +++++++++++++----- 5 files changed, 337 insertions(+), 276 deletions(-) delete mode 100644 migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py diff --git a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py b/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py deleted file mode 100644 index 306c772b..00000000 --- a/migrations/versions/394e7c39362d_add_token_count_to_queueitem_table.py +++ /dev/null @@ -1,134 +0,0 @@ -"""Add token_count to QueueItem table - -Revision ID: 394e7c39362d -Revises: 88b0fb10906f -Create Date: 2025-08-27 10:49:26.591473 - -""" - -import json -from collections.abc import Sequence - -import sqlalchemy as sa -from alembic import op - -from migrations.utils import column_exists -from src.config import settings - -# revision identifiers, used by Alembic. -revision: str = "394e7c39362d" -down_revision: str | None = "88b0fb10906f" -branch_labels: str | Sequence[str] | None = None -depends_on: str | Sequence[str] | None = None - -schema = settings.DB.SCHEMA - - -def upgrade() -> None: - op.add_column( - "queue", - sa.Column("token_count", sa.Integer(), nullable=False, server_default="0"), - schema=schema, - ) - - bind = op.get_bind() - - BATCH_SIZE = 500 # Process 500 items at a time - last_seen_id: int = 0 - - while True: - items_stmt = sa.text( - f""" - SELECT id, payload FROM {schema}.queue - WHERE processed = false - AND token_count = 0 - AND task_type IN ('representation', 'summary') - AND (payload->>'message_id') IS NOT NULL - AND id > :after_id - ORDER BY id - LIMIT :batch_size - """ - ).columns(id=sa.Integer, payload=sa.JSON) - - items_to_backfill = ( - bind.execute( - items_stmt, - {"batch_size": BATCH_SIZE, "after_id": last_seen_id}, - ) - .mappings() - .all() - ) - - if not items_to_backfill: - break # No more items to process - last_seen_id = items_to_backfill[-1]["id"] - - def _extract_message_id(payload: object) -> int | None: - if isinstance(payload, dict): - return payload.get("message_id") # pyright: ignore - if payload is None: - return None - try: - parsed = json.loads(str(payload)) - except (TypeError, json.JSONDecodeError): - return None - if isinstance(parsed, dict): - return parsed.get("message_id") # pyright: ignore - return None - - message_id_map = { - row["id"]: _extract_message_id(row["payload"]) for row in items_to_backfill - } - # Drop rows where message_id couldn't be extracted - message_id_map = { - queue_id: message_id - for queue_id, message_id in message_id_map.items() - if message_id is not None - } - - # Get token counts for the message IDs - token_ids = list(message_id_map.values()) - if not token_ids: - continue - - placeholders = ", ".join(f":id_{idx}" for idx in range(len(token_ids))) - token_counts_stmt = sa.text( - f""" - SELECT id, token_count FROM {schema}.messages - WHERE id IN ({placeholders}) - """ - ) - bind_params = {f"id_{idx}": token_id for idx, token_id in enumerate(token_ids)} - token_counts_result = bind.execute(token_counts_stmt, bind_params) - token_map = {msg_id: count for msg_id, count in token_counts_result.fetchall()} - - # Prepare parameters for the bulk update, skipping any queue items whose - # message_id was not found in the messages table. - update_params = [ - (qid, token_map.get(mid)) - for qid, mid in message_id_map.items() - if token_map.get(mid) is not None - ] - - if update_params: - update_stmt = sa.text( - f"UPDATE {schema}.queue SET token_count = :token_count WHERE id = :queue_id" - ) - - bind.execute( - update_stmt, - [ - {"queue_id": queue_id, "token_count": token_count} - for queue_id, token_count in update_params - ], - ) - - # If we fetched fewer items than the batch size, we are on the last batch - if len(items_to_backfill) < BATCH_SIZE: - break - - -def downgrade() -> None: - inspector = sa.inspect(op.get_bind()) - if column_exists("queue", "token_count", inspector): - op.drop_column("queue", "token_count", schema=schema) diff --git a/src/deriver/enqueue.py b/src/deriver/enqueue.py index 5974fc9e..c097e99c 100644 --- a/src/deriver/enqueue.py +++ b/src/deriver/enqueue.py @@ -165,7 +165,6 @@ def create_representation_record( "payload": processed_payload, "session_id": session_id, "task_type": "representation", - "token_count": message.get("token_count"), } @@ -198,7 +197,6 @@ def create_summary_record( "payload": processed_payload, "session_id": session_id, "task_type": "summary", - "token_count": message.get("token_count"), } diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 660165e3..f396efe3 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -8,7 +8,7 @@ from logging import getLogger import sentry_sdk from dotenv import load_dotenv from sentry_sdk.integrations.asyncio import AsyncioIntegration -from sqlalchemy import delete, select, update +from sqlalchemy import Integer, delete, select, update from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql import func @@ -234,50 +234,24 @@ class QueueManager: async with self.semaphore: message_count = 0 try: + from src.deriver.utils import parse_work_unit_key + + parsed_key = parse_work_unit_key(work_unit_key) + task_type = parsed_key["task_type"] + while not self.shutdown_event.is_set(): - candidate_messages = await self.get_message_batch( + messages_to_process: list[QueueItem] = await self.get_message_batch( work_unit_key, - limit=10, # hard limit of 10 messages per batch + task_type, ) - if not candidate_messages: - logger.debug(f"No more messages for work unit {work_unit_key}") - break - - 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: - msg_tokens = msg.token_count or 0 - # Always process at least one message, even if over limit - if ( - not messages_to_process - or token_count + msg_tokens <= max_tokens - ): - messages_to_process.append(msg) - token_count += msg_tokens - 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, - ) + logger.debug(f"No more messages for work unit {work_unit_key}") break # Process the batch/single item try: - raw_payloads = [msg.payload for msg in messages_to_process] - await process_items(task_type, raw_payloads) + payloads = [msg.payload for msg in messages_to_process] + await process_items(task_type, payloads) except Exception as e: logger.error( f"Error processing tasks for work unit {work_unit_key}: {e}", @@ -286,7 +260,6 @@ class QueueManager: if settings.SENTRY.ENABLED: sentry_sdk.capture_exception(e) - # Mark messages as processed (only for non-LLM errors) await self.mark_messages_as_processed( messages_to_process, work_unit_key ) @@ -343,19 +316,79 @@ class QueueManager: @sentry_sdk.trace async def get_message_batch( - self, work_unit_key: str, limit: int + self, work_unit_key: str, task_type: str ) -> list[QueueItem]: - """Get a batch of unprocessed messages for a specific work unit ordered by id.""" + """ + Get a batch of unprocessed messages for a specific work unit ordered by id. + For representation tasks, this will be a batch of messages up to REPRESENTATION_BATCH_MAX_TOKENS. + For other tasks, it will be a single message. + """ 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(limit) - ) - result = await db.execute(query) - messages = result.scalars().all() + if task_type != "representation": + # For non-representation tasks, just get the next single message. + query = ( + select(models.QueueItem) + .where(models.QueueItem.work_unit_key == work_unit_key) + .where(~models.QueueItem.processed) + .order_by(models.QueueItem.id) + .limit(1) + ) + result = await db.execute(query) + messages = result.scalars().all() + else: + # For representation tasks, get a batch based on token count. + # Always get at least the first message, then include additional messages + # as long as cumulative token count stays within limit. + # Join with messages table to get the actual token_count + + # Create CTE with row numbers and cumulative token counts + cte = ( + select( + models.QueueItem.id, + func.row_number() + .over(order_by=models.QueueItem.id) + .label("row_num"), + func.sum(models.Message.token_count) + .over(order_by=models.QueueItem.id) + .label("cumulative_token_count"), + ) + .select_from( + models.QueueItem.__table__.join( + models.Message.__table__, + func.cast( + models.QueueItem.payload["message_id"].astext, Integer + ) + == models.Message.id, + ) + ) + .where(models.QueueItem.work_unit_key == work_unit_key) + .where(~models.QueueItem.processed) + .order_by(models.QueueItem.id) + .cte() + ) + + # Select messages where either: + # 1. It's the first message (row_num = 1), OR + # 2. The cumulative token count is within the limit + query = ( + select(models.QueueItem) + .where( + models.QueueItem.id.in_( + select(cte.c.id).where( + (cte.c.row_num == 1) + | ( + cte.c.cumulative_token_count + <= settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + ) + ) + ) + ) + .order_by(models.QueueItem.id) + ) + + result = await db.execute(query) + 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() diff --git a/src/models.py b/src/models.py index 9b63d632..defe2110 100644 --- a/src/models.py +++ b/src/models.py @@ -376,7 +376,6 @@ class QueueItem(Base): task_type: Mapped[TaskType] = mapped_column(TEXT, nullable=False) payload: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) processed: Mapped[bool] = mapped_column(Boolean, default=False) - token_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False) def __repr__(self) -> str: return f"QueueItem(id={self.id}, session_id={self.session_id}, work_unit_key={self.work_unit_key}, task_type={self.task_type}, payload={self.payload}, processed={self.processed})" diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 8975d5f6..4f342261 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -118,22 +118,34 @@ class TestQueueProcessing: session, peers = sample_session_with_peers peer = peers[0] - payloads: list[Any] = [] - for i in range(3): - payloads.append( - create_queue_payload( # type: ignore[reportUnknownArgumentType] - message=models.Message( - id=i, - session_name=session.name, - workspace_name=session.workspace_name, - peer_name=peer.name, - content="hello", - ), # include id for payload builder - task_type="representation", - sender_name=peer.name, - target_name=peer.name, - ) + # Create and save messages to the database first + messages: list[models.Message] = [] + for _ in range(3): + message = models.Message( + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content="hello", + token_count=10, ) + db_session.add(message) + messages.append(message) + + await db_session.commit() + + # Refresh to get the actual IDs + for message in messages: + await db_session.refresh(message) + + payloads: list[Any] = [] + for message in messages: + payload = create_queue_payload( # type: ignore[reportUnknownArgumentType] + message=message, + task_type="representation", + sender_name=peer.name, + target_name=peer.name, + ) + payloads.append(payload) items = await add_queue_items(payloads, session.id) # Determine ascending order by DB id @@ -151,14 +163,20 @@ class TestQueueProcessing: first, second = ordered[0], ordered[1] qm = QueueManager() - batch = await qm.get_message_batch(first.work_unit_key, limit=1) + batch = await qm.get_message_batch( + first.work_unit_key, + task_type="representation", + ) 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() - batch2 = await qm.get_message_batch(first.work_unit_key, limit=1) + batch2 = await qm.get_message_batch( + first.work_unit_key, + task_type="representation", + ) nxt2 = batch2[0] if batch2 else None assert nxt2 is not None and nxt2.id == second.id @@ -253,17 +271,26 @@ class TestQueueProcessing: # Create messages with token counts that exceed batch limit limit = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS - messages = [ - models.Message( - id=i, + token_counts = [limit // 2, limit // 2, limit // 2] + + # Create and save messages to the database first + messages: list[models.Message] = [] + for i, token_count in enumerate(token_counts): + message = models.Message( 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([limit // 2, limit // 2, limit // 2]) - ] + db_session.add(message) + messages.append(message) + + await db_session.commit() + + # Refresh to get the actual IDs + for message in messages: + await db_session.refresh(message) # Create queue items with token counts payloads = [ @@ -280,7 +307,7 @@ class TestQueueProcessing: from src.deriver.utils import get_work_unit_key queue_items: list[models.QueueItem] = [] - for payload, message in zip(payloads, messages, strict=False): + for payload in payloads: task_type = payload.get("task_type", "unknown") work_unit_key = get_work_unit_key(task_type, payload) @@ -290,7 +317,6 @@ class TestQueueProcessing: 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) @@ -325,68 +351,6 @@ class TestQueueProcessing: assert processed_batches[1]["payload_count"] == 1 assert all(b["task_type"] == "representation" for b in processed_batches) - @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=5, # 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 - - # Mark the first batch as processed - await qm.mark_messages_as_processed(batch, queue_items[0].work_unit_key) - - # Get the next batch - should return remaining 5 items - next_batch = await qm.get_message_batch(queue_items[0].work_unit_key, limit=10) - assert len(next_batch) == 5 # Should return remaining items - @pytest.mark.asyncio async def test_single_message_processing( self, @@ -401,6 +365,7 @@ class TestQueueProcessing: peer = peers[0] # Create two summary messages + token_counts = [500, 600] messages = [ models.Message( id=999, @@ -408,7 +373,6 @@ class TestQueueProcessing: workspace_name=session.workspace_name, peer_name=peer.name, content="First summary message", - token_count=500, ), models.Message( id=1000, @@ -416,7 +380,6 @@ class TestQueueProcessing: workspace_name=session.workspace_name, peer_name=peer.name, content="Second summary message", - token_count=600, ), ] @@ -426,6 +389,7 @@ class TestQueueProcessing: payload = create_queue_payload( message, "summary", message_seq_in_session=i + 1 ) + payload["token_count"] = token_counts[i] from src.deriver.utils import get_work_unit_key work_unit_key = get_work_unit_key("summary", payload) @@ -436,7 +400,6 @@ class TestQueueProcessing: 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) @@ -490,5 +453,207 @@ class TestQueueProcessing: # Optionally verify the items have the expected token counts from the messages expected_token_counts = [500, 600] # From the test messages - actual_token_counts = [item.token_count for item in processed_items] + actual_token_counts = [ + item.payload.get("token_count") or 0 for item in processed_items + ] assert sorted(actual_token_counts) == sorted(expected_token_counts) + + @pytest.mark.asyncio + async def test_first_message_exceeds_token_limit_still_included( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + ) -> None: + """Test that if the first message exceeds BATCH_MAX_TOKENS, it's still included alone""" + from unittest.mock import patch + + session, peers = sample_session_with_peers + peer = peers[0] + + # Create messages where first message exceeds the batch limit + limit = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + token_counts = [limit + 1000, 100, 200] # First message way over limit + + # Create and save messages to the database first + messages: list[models.Message] = [] + for i, token_count in enumerate(token_counts): + message = models.Message( + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content=f"Test message {i}", + token_count=token_count, + ) + db_session.add(message) + messages.append(message) + + await db_session.commit() + + # Refresh to get the actual IDs + for message in messages: + await db_session.refresh(message) + + # Create queue items + 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 to queue + from src.deriver.utils import get_work_unit_key + + queue_items: list[models.QueueItem] = [] + for payload in payloads: + 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, + ) + 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: first large message alone, then second and third together + assert len(processed_batches) == 2 + assert ( + processed_batches[0]["payload_count"] == 1 + ) # First message (over limit) alone + assert processed_batches[1]["payload_count"] == 2 # Second and third messages + assert all(b["task_type"] == "representation" for b in processed_batches) + + @pytest.mark.asyncio + async def test_message_exactly_at_token_limit( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + ) -> None: + """Test boundary condition when cumulative sum exactly equals limit""" + from unittest.mock import patch + + session, peers = sample_session_with_peers + peer = peers[0] + + # Create messages that test the exact boundary + limit = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + token_counts = [ + limit // 2, + limit // 2, + 1, + ] # First two exactly at limit, third exceeds + + # Create and save messages to the database first + messages: list[models.Message] = [] + for i, token_count in enumerate(token_counts): + message = models.Message( + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content=f"Test message {i}", + token_count=token_count, + ) + db_session.add(message) + messages.append(message) + + await db_session.commit() + + # Refresh to get the actual IDs + for message in messages: + await db_session.refresh(message) + + # Create queue items + 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 to queue + from src.deriver.utils import get_work_unit_key + + queue_items: list[models.QueueItem] = [] + for payload in payloads: + 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, + ) + 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: first two messages together (exactly at limit), third alone + assert len(processed_batches) == 2 + assert ( + processed_batches[0]["payload_count"] == 2 + ) # First two messages (exactly at limit) + assert ( + processed_batches[1]["payload_count"] == 1 + ) # Third message (exceeds limit) + assert all(b["task_type"] == "representation" for b in processed_batches) From f1bdcf0f2efd901811650f469c1a5b22e92cd6f3 Mon Sep 17 00:00:00 2001 From: Benjamin McCormick Date: Wed, 24 Sep 2025 12:10:18 -0400 Subject: [PATCH 16/18] fix: langfuse sdk changes --- src/utils/embedding_store.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/utils/embedding_store.py b/src/utils/embedding_store.py index 08fd5bf4..8144b070 100644 --- a/src/utils/embedding_store.py +++ b/src/utils/embedding_store.py @@ -69,7 +69,7 @@ class EmbeddingStore: conclusions, similarity_threshold=similarity_threshold ) if settings.LANGFUSE_PUBLIC_KEY: - langfuse_context.update_current_observation( + lf.update_current_trace( input={"observations": [obs.model_dump() for obs in observations]}, output={"unique_conclusions": unique_conclusions}, ) From 70b6965939a2cdb94bc3340c96fe20f0b727eef3 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 24 Sep 2025 12:13:15 -0400 Subject: [PATCH 17/18] fix: CR comments --- src/deriver/queue_manager.py | 4 +--- src/routers/messages.py | 2 -- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index f396efe3..56ea38e4 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -16,6 +16,7 @@ from src import models from src.config import settings from src.dependencies import tracked_db from src.deriver.consumer import process_items +from src.deriver.utils import parse_work_unit_key from src.models import QueueItem logger = getLogger(__name__) @@ -234,8 +235,6 @@ class QueueManager: async with self.semaphore: message_count = 0 try: - from src.deriver.utils import parse_work_unit_key - parsed_key = parse_work_unit_key(work_unit_key) task_type = parsed_key["task_type"] @@ -281,7 +280,6 @@ class QueueManager: if removed and message_count > 0: # Only publish webhook if we actually removed an active session try: - from src.deriver.utils import parse_work_unit_key from src.webhooks.events import ( QueueEmptyEvent, publish_webhook_event, diff --git a/src/routers/messages.py b/src/routers/messages.py index 27d89a14..3c59dbb3 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -66,7 +66,6 @@ async def create_messages_for_session( "content": message.content, "peer_name": message.peer_name, "created_at": message.created_at, - "token_count": message.token_count, } for message in created_messages ] @@ -128,7 +127,6 @@ async def create_messages_with_file( "content": message.content, "peer_name": message.peer_name, "created_at": message.created_at, - "token_count": message.token_count, } for message in created_messages ] From 57977b541e98632bfe505a3b8fdcc93c21b8eaee Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 24 Sep 2025 13:22:29 -0400 Subject: [PATCH 18/18] fix: Integer -> BigInteger --- src/deriver/queue_manager.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index 56ea38e4..e17de306 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -8,7 +8,8 @@ from logging import getLogger import sentry_sdk from dotenv import load_dotenv from sentry_sdk.integrations.asyncio import AsyncioIntegration -from sqlalchemy import Integer, delete, select, update +from sqlalchemy import BigInteger, delete, select, update +from sqlalchemy.dialects.postgresql import insert from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.sql import func @@ -167,8 +168,6 @@ class QueueManager: async def claim_work_units( self, db: AsyncSession, work_unit_keys: Sequence[str] ) -> list[str]: - from sqlalchemy.dialects.postgresql import insert - values = [{"work_unit_key": key} for key in work_unit_keys] stmt = ( @@ -354,7 +353,8 @@ class QueueManager: models.QueueItem.__table__.join( models.Message.__table__, func.cast( - models.QueueItem.payload["message_id"].astext, Integer + models.QueueItem.payload["message_id"].astext, + BigInteger, ) == models.Message.id, )