From 110787cdca6f5e82a8e8a476b139abb6c6af7e4a Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 26 Jan 2026 18:00:53 -0500 Subject: [PATCH] feat: use messages from queue items for rep completed token count (#350) --- src/deriver/consumer.py | 5 +++-- src/deriver/deriver.py | 15 +++++++++------ src/deriver/queue_manager.py | 7 ++++++- tests/deriver/test_queue_processing.py | 6 +++--- tests/integration/test_token_metrics.py | 6 +++--- 5 files changed, 24 insertions(+), 15 deletions(-) diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 28eca6d8..9077fea4 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -158,7 +158,7 @@ async def process_representation_batch( *, observers: list[str] | None, observed: str | None, - queue_items_count: int, + queue_item_message_ids: list[int], ) -> None: """ Prepares and processes a batch of messages for representation tasks. @@ -168,6 +168,7 @@ async def process_representation_batch( message_level_configuration: Resolved configuration for this batch observers: List of observers for the messages observed: The observed of the messages + queue_item_message_ids: Message IDs from queue items """ if not messages or not messages[0]: logger.debug("process_representation_batch received no messages") @@ -181,7 +182,7 @@ async def process_representation_batch( message_level_configuration, observers=observers, observed=observed, - queue_items_count=queue_items_count, + queue_item_message_ids=queue_item_message_ids, ) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index bf786761..b8735e3c 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -20,7 +20,7 @@ from src.utils.clients import honcho_llm_call from src.utils.config_helpers import get_configuration from src.utils.formatting import format_new_turn_with_timestamp from src.utils.representation import PromptRepresentation, Representation -from src.utils.tokens import estimate_tokens, track_deriver_input_tokens +from src.utils.tokens import track_deriver_input_tokens from .prompts import estimate_minimal_deriver_prompt_tokens, minimal_deriver_prompt @@ -34,7 +34,7 @@ async def process_representation_tasks_batch( *, observers: list[str], observed: str, - queue_items_count: int, + queue_item_message_ids: list[int], ) -> None: """ Process messages with minimal overhead - single LLM call, save to multiple collections. @@ -44,7 +44,7 @@ async def process_representation_tasks_batch( message_level_configuration: Optional configuration override. observers: List of observer peer IDs (collections to save to). observed: The observed peer ID. - queue_items_count: Number of QueueItem records being processed in this batch. + queue_item_message_ids: Message IDs from queue items being processed """ if not messages: return @@ -93,9 +93,12 @@ async def process_representation_tasks_batch( for msg in messages ) - # Track token usage + # Track token usage - count only tokens from messages being processed prompt_tokens = estimate_minimal_deriver_prompt_tokens() - messages_tokens = estimate_tokens(formatted_messages) + queue_item_message_ids_set = set(queue_item_message_ids) + messages_tokens = sum( + msg.token_count for msg in messages if msg.id in queue_item_message_ids_set + ) track_deriver_input_tokens( task_type=DeriverTaskTypes.INGESTION, components={ @@ -235,7 +238,7 @@ async def process_representation_tasks_batch( workspace_name=latest_message.workspace_name, session_name=latest_message.session_name, observed=observed, - queue_items_processed=queue_items_count, + queue_items_processed=len(queue_item_message_ids), earliest_message_id=earliest_message.public_id, latest_message_id=latest_message.public_id, message_count=len(messages), diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index c9fd2893..c1f48d86 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -450,12 +450,17 @@ class QueueManager: else: observers = [] + queue_item_message_ids = [ + item.message_id + for item in items_to_process + if item.message_id is not None + ] await process_representation_batch( messages_context, message_level_configuration, observers=observers, observed=work_unit.observed, - queue_items_count=len(items_to_process), + queue_item_message_ids=queue_item_message_ids, ) await self.mark_queue_items_as_processed( items_to_process, work_unit_key diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 9ca78f97..51c8af07 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -347,7 +347,7 @@ class TestQueueProcessing: *, observed: str | None = None, # pyright: ignore[reportUnusedParameter] observers: list[str] | None = None, # pyright: ignore[reportUnusedParameter] - queue_items_count: int | None = None, # pyright: ignore[reportUnusedParameter] + queue_item_message_ids: list[int] | None = None, # pyright: ignore[reportUnusedParameter] ) -> None: processed_batches.append( { @@ -910,7 +910,7 @@ class TestQueueProcessing: *, observed: str | None = None, # pyright: ignore[reportUnusedParameter] observers: list[str] | None = None, # pyright: ignore[reportUnusedParameter] - queue_items_count: int | None = None, # pyright: ignore[reportUnusedParameter] + queue_item_message_ids: list[int] | None = None, # pyright: ignore[reportUnusedParameter] ) -> None: processed_batches.append( { @@ -1029,7 +1029,7 @@ class TestQueueProcessing: *, observed: str | None = None, # pyright: ignore[reportUnusedParameter] observers: list[str] | None = None, # pyright: ignore[reportUnusedParameter] - queue_items_count: int | None = None, # pyright: ignore[reportUnusedParameter] + queue_item_message_ids: list[int] | None = None, # pyright: ignore[reportUnusedParameter] ) -> None: processed_batches.append( { diff --git a/tests/integration/test_token_metrics.py b/tests/integration/test_token_metrics.py index 88bfa369..31163474 100644 --- a/tests/integration/test_token_metrics.py +++ b/tests/integration/test_token_metrics.py @@ -261,7 +261,7 @@ class TestDeriverIngestionMetrics: message_level_configuration=create_test_configuration(), observers=[peer.name], observed=peer.name, - queue_items_count=len(messages), + queue_item_message_ids=[m.id for m in messages], ) # Verify output tokens metric @@ -318,7 +318,7 @@ class TestDeriverIngestionMetrics: message_level_configuration=create_test_configuration(), observers=[peer.name], observed=peer.name, - queue_items_count=len(messages), + queue_item_message_ids=[m.id for m in messages], ) metric_checker.assert_delta( @@ -375,7 +375,7 @@ class TestDeriverIngestionMetrics: message_level_configuration=create_test_configuration(), observers=[peer.name], observed=peer.name, - queue_items_count=len(messages), + queue_item_message_ids=[m.id for m in messages], ) # Verify messages tokens were tracked (should be > 0)