fix: CR comments
This commit is contained in:
parent
dfa7577eb8
commit
5301914847
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue