diff --git a/.env.template b/.env.template index abb4e5b3..9cc57de0 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 # Bytes +# MAX_MESSAGE_SIZE=25000 # Characters # 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) diff --git a/config.toml.example b/config.toml.example index 98d84f5a..3be89786 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 # Characters 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 diff --git a/src/config.py b/src/config.py index 7d91d19c..fec3b882 100644 --- a/src/config.py +++ b/src/config.py @@ -212,6 +212,22 @@ class DeriverSettings(HonchoSettings): int, Field(default=100, gt=0, le=500) ] = 100 + REPRESENTATION_BATCH_MAX_TOKENS: Annotated[ + int, + Field( + default=4096, + ge=1, + ), + ] = 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/__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 015763c6..c564722f 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, ) @@ -27,25 +28,36 @@ console = Console(markup=True) lf = get_client() -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) + if not queue_payloads or not queue_payloads[0]: + 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), + 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: lf.update_current_trace( # type: ignore @@ -54,13 +66,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: lf.update_current_trace( @@ -68,17 +83,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 ec7fb48f..9b53d3e6 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -57,7 +57,7 @@ async def critical_analysis_call( message_created_at: datetime.datetime, working_representation: str | None, history: str, - new_turn: str, + new_turns: list[str], ) -> ReasoningResponse: prompt = critical_analysis_prompt( peer_id=peer_id, @@ -65,7 +65,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, ) response = await honcho_llm_call( @@ -114,27 +114,37 @@ async def peer_card_call( return response.content -@conditional_observe @sentry_sdk.trace -async def process_representation_task( - payload: RepresentationPayload, +async def process_representation_tasks_batch( + payloads: list[RepresentationPayload], ) -> None: """ - Process a representation task by extracting insights and updating working representations. + Process a batch of representation tasks by extracting insights and updating working representations. """ + if not payloads or len(payloads) == 0: + return + + payloads.sort(key=lambda x: x.message_id) + + latest_payload = payloads[-1] + earliest_payload = payloads[0] + # 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, + latest_payload.workspace_name, + latest_payload.session_name, token_limit=settings.DERIVER.CONTEXT_TOKEN_LIMIT, - cutoff=payload.message_id, + cutoff=earliest_payload.message_id, include_summary=True, ) @@ -144,9 +154,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 ) @@ -154,21 +164,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: @@ -176,10 +186,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 @@ -212,8 +222,12 @@ async def process_representation_task( ) else: # No existing working representation, use global search + query_text = [payload.content for payload in payloads] + query_text = "\n".join( + query_text + ) # TODO: consider a smarter strategy than concatenation working_representation = await embedding_store.get_relevant_observations( - query=payload.content, + query=query_text, conversation_context=formatted_history, for_reasoning=True, ) @@ -224,7 +238,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", @@ -237,10 +251,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) @@ -260,12 +277,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", @@ -274,13 +290,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, - "", + "count", ) 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: @@ -289,14 +305,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 @@ -312,47 +335,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: lf.update_current_generation( 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 @@ -422,6 +447,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 @@ -436,7 +462,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", @@ -447,13 +473,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", @@ -472,7 +498,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,8 +510,10 @@ 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.""" # Use the utility function to find new observations @@ -533,9 +561,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") @@ -569,9 +597,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/enqueue.py b/src/deriver/enqueue.py index 00028cda..c097e99c 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 diff --git a/src/deriver/prompts.py b/src/deriver/prompts.py index c6f37603..30009fbd 100644 --- a/src/deriver/prompts.py +++ b/src/deriver/prompts.py @@ -15,7 +15,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. @@ -26,7 +26,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 @@ -55,6 +55,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. @@ -91,10 +93,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/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index f6de3e4d..e17de306 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -8,17 +8,18 @@ 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 BigInteger, delete, select, update +from sqlalchemy.dialects.postgresql import insert 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.deriver.utils import parse_work_unit_key from src.models import QueueItem -from .. import models -from ..dependencies import tracked_db -from .consumer import process_item - logger = getLogger(__name__) load_dotenv(override=True) @@ -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 = ( @@ -232,52 +231,39 @@ 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: + parsed_key = parse_work_unit_key(work_unit_key) + task_type = parsed_key["task_type"] + while not self.shutdown_event.is_set(): - message = await self.get_next_message(work_unit_key) - if not message: + messages_to_process: list[QueueItem] = await self.get_message_batch( + work_unit_key, + task_type, + ) + if not messages_to_process: logger.debug(f"No more messages for work unit {work_unit_key}") break - message_count += 1 + # 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}" - ) + payloads = [msg.payload for msg in messages_to_process] + await process_items(task_type, payloads) 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() + 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", @@ -285,9 +271,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") @@ -296,7 +279,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, @@ -330,22 +312,104 @@ 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: - 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) - message = result.scalar_one_or_none() + async def get_message_batch( + self, work_unit_key: str, task_type: str + ) -> list[QueueItem]: + """ + 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: + 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, + BigInteger, + ) + == 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() - 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/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}, ) diff --git a/tests/deriver/test_deriver_processing.py b/tests/deriver/test_deriver_processing.py index deaa085a..020d9429 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,92 @@ 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 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", + 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] diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index e38571c0..4f342261 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 @@ -117,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 @@ -150,13 +163,21 @@ 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, + 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() - nxt2 = await qm.get_next_message(first.work_unit_key) + 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 @pytest.mark.asyncio @@ -234,3 +255,405 @@ 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 + limit = settings.DERIVER.REPRESENTATION_BATCH_MAX_TOKENS + 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, + ) + 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 = [ + 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 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 due to token limits + assert len(processed_batches) == 2 + 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_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 + token_counts = [500, 600] + messages = [ + models.Message( + id=999, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content="First summary message", + ), + models.Message( + id=1000, + session_name=session.name, + workspace_name=session.workspace_name, + peer_name=peer.name, + content="Second summary message", + ), + ] + + # 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 + ) + payload["token_count"] = token_counts[i] + 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, + ) + 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) + + # 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.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) 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