fix: CR comments

This commit is contained in:
Rajat Ahuja 2025-09-19 12:24:01 -04:00
parent dfa7577eb8
commit 5301914847
5 changed files with 47 additions and 9 deletions

View File

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

View File

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

View File

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

View File

@ -12,13 +12,12 @@ from sqlalchemy import delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql import func
from src import 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

View File

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