From 5d992bc65afcfbc05a5911ab4edbaa88ef64c690 Mon Sep 17 00:00:00 2001 From: Ulysse Pence <736903+ulyssepence@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:42:48 -0400 Subject: [PATCH] feat(api): Export deriver backlog as metrics from API endpoint (#1115) --- src/backlog.py | 146 +++++++++ src/config.py | 3 + src/crud/__init__.py | 7 +- src/crud/deriver.py | 152 ++++++++- src/deriver/queue_manager.py | 45 +-- src/dreamer/dream_due.py | 216 +++++++++++++ src/main.py | 12 + src/reconciler/embed_now.py | 4 +- src/reconciler/sync_vectors.py | 8 +- src/routers/deriver_metrics.py | 41 +++ src/schemas/__init__.py | 2 + src/schemas/internal.py | 11 + src/telemetry/prometheus/metrics.py | 107 ++++++ tests/crud/test_deriver_metrics_query.py | 394 +++++++++++++++++++++++ tests/dreamer/test_dream_due.py | 321 ++++++++++++++++++ tests/telemetry/test_metric_zero_init.py | 18 ++ tests/test_deriver_metrics.py | 207 ++++++++++++ 17 files changed, 1656 insertions(+), 38 deletions(-) create mode 100644 src/backlog.py create mode 100644 src/dreamer/dream_due.py create mode 100644 src/routers/deriver_metrics.py create mode 100644 tests/crud/test_deriver_metrics_query.py create mode 100644 tests/dreamer/test_dream_due.py create mode 100644 tests/test_deriver_metrics.py diff --git a/src/backlog.py b/src/backlog.py new file mode 100644 index 00000000..1180be57 --- /dev/null +++ b/src/backlog.py @@ -0,0 +1,146 @@ +"""Read-only polling of the deriver's outstanding work. Schedules nothing.""" + +import asyncio +import contextlib +import time +from dataclasses import dataclass, field +from logging import getLogger + +import sentry_sdk + +from src import crud, schemas +from src.config import settings +from src.dependencies import tracked_db +from src.dreamer.dream_due import count_due_dreams +from src.telemetry import prometheus_metrics + +logger = getLogger(__name__) + + +def active_work_seconds() -> float: + """The value reported when work is ready for a deriver now.""" + return float(max(settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS, 1)) + + +@dataclass +class DeriverMetricsSnapshot: + """The last good poll result, served to callers of the route.""" + + signal_seconds: float = 0.0 + dreams_due: int = 0 + stats: schemas.DeriverMetrics = field(default_factory=schemas.DeriverMetrics) + measured_at: float | None = None + + @property + def age_seconds(self) -> float | None: + if self.measured_at is None: + return None + return max(0.0, time.time() - self.measured_at) + + +def outstanding_work_seconds( + stats: schemas.DeriverMetrics, *, dreams_due: int +) -> float: + """Seconds of outstanding deriver work, 0 when there is nothing to do.""" + if ( + stats.eligible_work_units > 0 + or stats.claimed_work_units > 0 + or stats.embeddings_pending_due > 0 + or dreams_due > 0 + ): + return active_work_seconds() + if stats.pending_items > 0: + return stats.oldest_pending_age_seconds + return 0.0 + + +class DeriverMetricsPoller: + """Refreshes the deriver gauges and the cached snapshot on a timer.""" + + def __init__(self) -> None: + self._task: asyncio.Task[None] | None = None + self._shutdown_event: asyncio.Event = asyncio.Event() + self._snapshot: DeriverMetricsSnapshot = DeriverMetricsSnapshot() + self._next_dream_poll: float | None = None + self._dreams_due: int = 0 + + @property + def snapshot(self) -> DeriverMetricsSnapshot: + return self._snapshot + + async def start(self) -> None: + if self._task is not None: + logger.warning("DeriverMetricsPoller already running") + return + self._shutdown_event.clear() + self._task = asyncio.create_task(self._loop()) + logger.info( + "DeriverMetricsPoller started, interval %ss", + settings.DERIVER.BACKLOG_METRICS_POLL_INTERVAL_SECONDS, + ) + + async def shutdown(self) -> None: + if self._task is None: + return + logger.info("Shutting down DeriverMetricsPoller...") + self._shutdown_event.set() + try: + await asyncio.wait_for(self._task, timeout=5.0) + except TimeoutError: + logger.warning("DeriverMetricsPoller shutdown timed out, cancelling task") + self._task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await self._task + self._task = None + logger.info("DeriverMetricsPoller stopped") + + async def _loop(self) -> None: + interval = settings.DERIVER.BACKLOG_METRICS_POLL_INTERVAL_SECONDS + while not self._shutdown_event.is_set(): + try: + await self.refresh() + except Exception as e: + logger.error("DeriverMetricsPoller refresh failed: %s", e) + if settings.SENTRY.ENABLED: + sentry_sdk.capture_exception(e) + with contextlib.suppress(TimeoutError): + await asyncio.wait_for(self._shutdown_event.wait(), timeout=interval) + + async def refresh(self) -> None: + """One read-only pass. The snapshot only advances on a complete pass.""" + async with tracked_db("deriver_metrics", read_only=True) as db: + stats = await crud.get_deriver_metrics(db) + if self._dream_poll_due(): + self._dreams_due = await count_due_dreams(db) + self._next_dream_poll = ( + time.monotonic() + settings.DREAM.DUE_POLL_INTERVAL_SECONDS + ) + + signal = outstanding_work_seconds(stats, dreams_due=self._dreams_due) + measured_at = time.time() + + self._snapshot = DeriverMetricsSnapshot( + signal_seconds=signal, + dreams_due=self._dreams_due, + stats=stats, + measured_at=measured_at, + ) + + metrics = prometheus_metrics + metrics.set_deriver_metrics( + eligible_work_units=stats.eligible_work_units, + claimed_work_units=stats.claimed_work_units, + pending_items=stats.pending_items, + oldest_pending_age_seconds=stats.oldest_pending_age_seconds, + embeddings_pending=stats.embeddings_pending, + embeddings_pending_due=stats.embeddings_pending_due, + ) + metrics.set_dreams_due(count=self._dreams_due) + metrics.set_deriver_outstanding_work(seconds=signal) + metrics.set_deriver_metrics_last_success(timestamp=measured_at) + + def _dream_poll_due(self) -> bool: + """The dream query is far more expensive, so it runs on its own spacing.""" + return ( + self._next_dream_poll is None or time.monotonic() >= self._next_dream_poll + ) diff --git a/src/config.py b/src/config.py index 993f9bfb..80827327 100644 --- a/src/config.py +++ b/src/config.py @@ -972,6 +972,8 @@ class DeriverSettings(HonchoSettings): # When enabled, bypasses the batch token threshold and processes work immediately FLUSH_ENABLED: bool = False + BACKLOG_METRICS_POLL_INTERVAL_SECONDS: Annotated[int, Field(default=30, ge=1)] = 30 + @model_validator(mode="before") @classmethod def _merge_model_config_defaults(cls, data: Any) -> Any: @@ -1351,6 +1353,7 @@ class DreamSettings(HonchoSettings): DOCUMENT_THRESHOLD: Annotated[int, Field(default=50, gt=0, le=1000)] = 50 IDLE_TIMEOUT_MINUTES: Annotated[int, Field(default=60, gt=0, le=1440)] = 60 MIN_HOURS_BETWEEN_DREAMS: Annotated[int, Field(default=8, gt=0, le=72)] = 8 + DUE_POLL_INTERVAL_SECONDS: Annotated[int, Field(default=300, ge=1)] = 300 ENABLED_TYPES: list[str] = ["omni"] # Agent iteration limit - increased for extended reasoning workflow diff --git a/src/crud/__init__.py b/src/crud/__init__.py index 0e920717..ac17af3f 100644 --- a/src/crud/__init__.py +++ b/src/crud/__init__.py @@ -3,7 +3,11 @@ from .collection import ( get_or_create_collection, update_collection_internal_metadata, ) -from .deriver import get_deriver_status, get_queue_status +from .deriver import ( + get_deriver_metrics, + get_deriver_status, + get_queue_status, +) from .document import ( CreateDocumentsResult, create_documents, @@ -105,6 +109,7 @@ __all__ = [ "get_or_create_collection", "update_collection_internal_metadata", # Deriver + "get_deriver_metrics", "get_deriver_status", "get_queue_status", # Document diff --git a/src/crud/deriver.py b/src/crud/deriver.py index 0852a479..770ba929 100644 --- a/src/crud/deriver.py +++ b/src/crud/deriver.py @@ -1,15 +1,165 @@ from collections.abc import Sequence +from datetime import UTC, datetime, timedelta from logging import getLogger from typing import Any -from sqlalchemy import Select, case, func, or_, select +from sqlalchemy import ColumnElement, Select, case, func, or_, select from sqlalchemy.engine import Row from sqlalchemy.ext.asyncio import AsyncSession from src import models, schemas +from src.config import settings logger = getLogger(__name__) +REPRESENTATION_WORK_UNIT_PREFIX = "representation:" + + +def representation_batch_threshold_clause( + *, + work_unit_key: ColumnElement[str], + total_tokens: ColumnElement[Any], + oldest_created_at: ColumnElement[Any], +) -> ColumnElement[bool] | None: + """The batch gate a representation work unit passes before it is claimable, or None when no gate applies.""" + if settings.DERIVER.FLUSH_ENABLED: + return None + + target_tokens = settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS + if target_tokens <= 0: + return None + + threshold: ColumnElement[bool] = func.coalesce(total_tokens, 0) >= target_tokens + + max_age_seconds = settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS + if max_age_seconds > 0: + threshold = or_( + threshold, + oldest_created_at <= func.now() - timedelta(seconds=max_age_seconds), + ) + + return or_( + ~work_unit_key.startswith(REPRESENTATION_WORK_UNIT_PREFIX), + threshold, + ) + + +def unclaimed_work_unit_clause( + work_unit_key: ColumnElement[str], +) -> ColumnElement[bool]: + """No claim row exists for this work unit, stale ones included.""" + return ( + ~select(models.ActiveQueueSession.id) + .where(models.ActiveQueueSession.work_unit_key == work_unit_key) + .exists() + ) + + +def stale_claim_cutoff() -> datetime: + return datetime.now(UTC) - timedelta( + minutes=settings.DERIVER.STALE_SESSION_TIMEOUT_MINUTES + ) + + +def not_live_claimed_work_unit_clause( + work_unit_key: ColumnElement[str], +) -> ColumnElement[bool]: + """No claim refreshed inside the stale timeout exists, so a stale claim leaves its work unit claimable.""" + return ( + ~select(models.ActiveQueueSession.id) + .where( + models.ActiveQueueSession.work_unit_key == work_unit_key, + models.ActiveQueueSession.last_updated >= stale_claim_cutoff(), + ) + .exists() + ) + + +async def get_deriver_metrics(db: AsyncSession) -> schemas.DeriverMetrics: + """Count the outstanding deriver work in the whole database, read-only.""" + from src.reconciler.sync_vectors import backoff_eligible # noqa: PLC0415 + + token_stats = ( + select( + models.QueueItem.work_unit_key, + func.sum(models.Message.token_count).label("total_tokens"), + func.min(models.QueueItem.created_at).label("oldest_created_at"), + ) + .join(models.Message, models.QueueItem.message_id == models.Message.id) + .where(~models.QueueItem.processed) + .where( + models.QueueItem.work_unit_key.startswith(REPRESENTATION_WORK_UNIT_PREFIX) + ) + .group_by(models.QueueItem.work_unit_key) + .subquery() + ) + + work_units = ( + select(models.QueueItem.work_unit_key) + .where(~models.QueueItem.processed) + .group_by(models.QueueItem.work_unit_key) + .subquery() + ) + + eligible = ( + select(func.count()) + .select_from(work_units) + .outerjoin( + token_stats, + work_units.c.work_unit_key == token_stats.c.work_unit_key, + ) + .where(not_live_claimed_work_unit_clause(work_units.c.work_unit_key)) + ) + + threshold_clause = representation_batch_threshold_clause( + work_unit_key=work_units.c.work_unit_key, + total_tokens=token_stats.c.total_tokens, + oldest_created_at=token_stats.c.oldest_created_at, + ) + if threshold_clause is not None: + eligible = eligible.where(threshold_clause) + + claimed = ( + select(func.count()) + .select_from(models.ActiveQueueSession) + .where(models.ActiveQueueSession.last_updated >= stale_claim_cutoff()) + ) + + pending = select( + func.count(models.QueueItem.id), + func.coalesce( + func.extract("epoch", func.now() - func.min(models.QueueItem.created_at)), + 0, + ), + ).where(~models.QueueItem.processed) + + embeddings = select( + func.count(), + func.coalesce( + func.sum( + case( + (backoff_eligible(models.MessageEmbedding.last_sync_at), 1), + else_=0, + ) + ), + 0, + ), + ).where(models.MessageEmbedding.sync_state == "pending") + + eligible_count = (await db.execute(eligible)).scalar_one() + claimed_count = (await db.execute(claimed)).scalar_one() + pending_count, oldest_age = (await db.execute(pending)).one() + embeddings_pending, embeddings_due = (await db.execute(embeddings)).one() + + return schemas.DeriverMetrics( + eligible_work_units=int(eligible_count), + claimed_work_units=int(claimed_count), + pending_items=int(pending_count), + oldest_pending_age_seconds=float(oldest_age), + embeddings_pending=int(embeddings_pending), + embeddings_pending_due=int(embeddings_due), + ) + async def get_queue_status( db: AsyncSession, diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index b98c0ef6..493cddfa 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -15,7 +15,7 @@ from dotenv import load_dotenv from nanoid import generate as generate_nanoid from sentry_sdk.integrations.asyncio import AsyncioIntegration from sentry_sdk.integrations.sqlalchemy import SqlalchemyIntegration -from sqlalchemy import Text, and_, delete, literal, or_, select, update +from sqlalchemy import Text, and_, delete, literal, select, update from sqlalchemy.dialects.postgresql import insert from sqlalchemy.engine import CursorResult from sqlalchemy.ext.asyncio import AsyncSession @@ -24,6 +24,11 @@ from sqlalchemy.sql import func from src import models from src.cache.client import close_cache, init_cache from src.config import settings +from src.crud.deriver import ( + REPRESENTATION_WORK_UNIT_PREFIX, + representation_batch_threshold_clause, + unclaimed_work_unit_clause, +) from src.dependencies import tracked_db from src.deriver.consumer import ( process_item, @@ -353,7 +358,7 @@ class QueueManager: ) async with tracked_db("get_available_work_units") as db: - representation_prefix = "representation:" + representation_prefix = REPRESENTATION_WORK_UNIT_PREFIX token_stats_subq = ( select( models.QueueItem.work_unit_key, @@ -390,14 +395,7 @@ class QueueManager: token_stats_subq, work_units_subq.c.work_unit_key == token_stats_subq.c.work_unit_key, ) - .where( - ~select(models.ActiveQueueSession.id) - .where( - models.ActiveQueueSession.work_unit_key - == work_units_subq.c.work_unit_key - ) - .exists() - ) + .where(unclaimed_work_unit_clause(work_units_subq.c.work_unit_key)) .order_by( work_units_subq.c.oldest_created_at.asc(), work_units_subq.c.work_unit_key.asc(), @@ -406,26 +404,13 @@ class QueueManager: ) # Apply batch threshold filter (skip if FLUSH_ENABLED is True) - if not settings.DERIVER.FLUSH_ENABLED and work_unit_target_tokens > 0: - max_age_seconds = settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS - threshold_clause = ( - func.coalesce(token_stats_subq.c.total_tokens, 0) - >= work_unit_target_tokens - ) - if max_age_seconds > 0: - threshold_clause = or_( - threshold_clause, - token_stats_subq.c.oldest_created_at - <= func.now() - timedelta(seconds=max_age_seconds), - ) - query = query.where( - or_( - ~work_units_subq.c.work_unit_key.startswith( - representation_prefix - ), - threshold_clause, - ) - ) + threshold_clause = representation_batch_threshold_clause( + work_unit_key=work_units_subq.c.work_unit_key, + total_tokens=token_stats_subq.c.total_tokens, + oldest_created_at=token_stats_subq.c.oldest_created_at, + ) + if threshold_clause is not None: + query = query.where(threshold_clause) result = await db.execute(query) available_rows = result.all() diff --git a/src/dreamer/dream_due.py b/src/dreamer/dream_due.py new file mode 100644 index 00000000..08b68e2c --- /dev/null +++ b/src/dreamer/dream_due.py @@ -0,0 +1,216 @@ +"""Read-only count of the collections whose next dream is due. Enqueues nothing.""" + +from datetime import UTC, datetime, timedelta +from logging import getLogger +from typing import Any, cast + +from sqlalchemy import func, select +from sqlalchemy.dialects.postgresql import aggregate_order_by +from sqlalchemy.ext.asyncio import AsyncSession + +from src import models +from src.config import settings +from src.schemas import DreamType +from src.utils.config_helpers import get_configuration +from src.utils.work_unit import construct_work_unit_key + +logger = getLogger(__name__) + + +async def count_due_dreams(db: AsyncSession) -> int: + """Count collections past the threshold, the idle timeout, the min-hours gate, any earlier attempt, and the session's dream setting.""" + dream_types = [ + DreamType(dream_type) + for dream_type in settings.DREAM.ENABLED_TYPES + if dream_type == DreamType.OMNI.value + ] + if not settings.DREAM.ENABLED or not dream_types: + return 0 + + explicit_counts = ( + select( + models.Document.workspace_name, + models.Document.observer, + models.Document.observed, + func.count(models.Document.id).label("explicit_count"), + func.max(models.Document.created_at).label("newest_created_at"), + func.array_agg( + aggregate_order_by( + models.Document.session_name, models.Document.created_at.desc() + ) + )[1].label("newest_session_name"), + ) + .where(models.Document.level == "explicit") + .group_by( + models.Document.workspace_name, + models.Document.observer, + models.Document.observed, + ) + .subquery() + ) + + rows = ( + await db.execute( + select( + models.Collection.workspace_name, + models.Collection.observer, + models.Collection.observed, + models.Collection.internal_metadata, + func.coalesce(explicit_counts.c.explicit_count, 0), + explicit_counts.c.newest_created_at, + explicit_counts.c.newest_session_name, + ).outerjoin( + explicit_counts, + (models.Collection.workspace_name == explicit_counts.c.workspace_name) + & (models.Collection.observer == explicit_counts.c.observer) + & (models.Collection.observed == explicit_counts.c.observed), + ) + ) + ).all() + + now = datetime.now(UTC) + idle_cutoff = now - timedelta(minutes=settings.DREAM.IDLE_TIMEOUT_MINUTES) + candidates: dict[str, tuple[str, str, datetime]] = {} + + for row in rows: + workspace_name = cast(str, row[0]) + observer = cast(str, row[1]) + observed = cast(str, row[2]) + internal_metadata = cast("dict[str, Any] | None", row[3]) + explicit_count = cast(int, row[4]) + newest_created_at = cast("datetime | None", row[5]) + newest_session_name = cast("str | None", row[6]) + + dream_metadata: dict[str, Any] = (internal_metadata or {}).get("dream", {}) + since_last_dream = explicit_count - int( + dream_metadata.get("last_dream_document_count", 0) + ) + if since_last_dream < settings.DREAM.DOCUMENT_THRESHOLD: + continue + + if newest_created_at is None or newest_created_at > idle_cutoff: + continue + + if newest_session_name is None: + continue + + last_dream_at = cast("str | None", dream_metadata.get("last_dream_at")) + if last_dream_at and _within_min_hours_gate(last_dream_at, now): + continue + + for dream_type in dream_types: + work_unit_key = construct_work_unit_key( + workspace_name, + { + "task_type": "dream", + "observer": observer, + "observed": observed, + "dream_type": dream_type.value, + }, + ) + candidates[work_unit_key] = ( + workspace_name, + newest_session_name, + newest_created_at, + ) + + if not candidates: + return 0 + + attempt_rows = ( + await db.execute( + select( + models.QueueItem.work_unit_key, + func.max(models.QueueItem.created_at), + ) + .where( + models.QueueItem.task_type == "dream", + models.QueueItem.work_unit_key.in_(candidates.keys()), + ) + .group_by(models.QueueItem.work_unit_key) + ) + ).all() + newest_attempts: dict[str, datetime] = { + cast(str, row[0]): cast(datetime, row[1]) for row in attempt_rows + } + + unattempted = [ + (workspace_name, session_name) + for work_unit_key, ( + workspace_name, + session_name, + newest_created_at, + ) in candidates.items() + if work_unit_key not in newest_attempts + or newest_attempts[work_unit_key] < newest_created_at + ] + if not unattempted: + return 0 + + return await _count_with_dreams_enabled(db, unattempted) + + +async def _count_with_dreams_enabled( + db: AsyncSession, candidates: list[tuple[str, str]] +) -> int: + """Drop candidates whose resolved configuration has dreams turned off.""" + workspace_names = {workspace_name for workspace_name, _ in candidates} + session_keys = set(candidates) + + workspaces = { + workspace.name: workspace + for workspace in ( + await db.execute( + select(models.Workspace).where( + models.Workspace.name.in_(workspace_names) + ) + ) + ) + .scalars() + .all() + } + + sessions: dict[tuple[str, str], models.Session] = {} + if session_keys: + session_rows = ( + ( + await db.execute( + select(models.Session).where( + models.Session.workspace_name.in_(workspace_names), + models.Session.name.in_( + {session_name for _, session_name in candidates} + ), + ) + ) + ) + .scalars() + .all() + ) + sessions = { + (session.workspace_name, session.name): session for session in session_rows + } + + enabled = 0 + for workspace_name, session_name in candidates: + configuration = get_configuration( + None, + sessions.get((workspace_name, session_name)), + workspaces.get(workspace_name), + ) + if configuration.dream.enabled: + enabled += 1 + return enabled + + +def _within_min_hours_gate(last_dream_at: str, now: datetime) -> bool: + """True when the last dream is too recent for another one.""" + try: + last_dream_time = datetime.fromisoformat(last_dream_at) + except (ValueError, TypeError): + return False + + if last_dream_time.tzinfo is None: + last_dream_time = last_dream_time.replace(tzinfo=UTC) + + hours_since = (now - last_dream_time).total_seconds() / 3600 + return hours_since < settings.DREAM.MIN_HOURS_BETWEEN_DREAMS diff --git a/src/main.py b/src/main.py index a1ec9765..9a1d7e64 100644 --- a/src/main.py +++ b/src/main.py @@ -15,6 +15,7 @@ from sentry_sdk.integrations.sqlalchemy import SqlalchemyIntegration from sentry_sdk.integrations.starlette import StarletteIntegration from src._version import HONCHO_VERSION +from src.backlog import DeriverMetricsPoller from src.cache.client import close_cache, init_cache from src.config import settings from src.db import ( @@ -26,6 +27,7 @@ from src.db import ( from src.exceptions import HonchoException from src.routers import ( conclusions, + deriver_metrics, keys, messages, peers, @@ -135,12 +137,21 @@ async def lifespan(_: FastAPI): "Error initializing cache in api process; proceeding without cache: %s", e ) + deriver_metrics_poller = DeriverMetricsPoller() + deriver_metrics.set_deriver_metrics_poller(deriver_metrics_poller) + try: + await deriver_metrics_poller.start() + except Exception as e: + logger.error("Failed to start backlog metrics poller: %s", e) + try: yield finally: # Import here to avoid circular import at module load time from src.vector_store import close_external_vector_store + await deriver_metrics_poller.shutdown() + deriver_metrics.set_deriver_metrics_poller(None) await close_external_vector_store() await close_cache() await engine.dispose() @@ -189,6 +200,7 @@ app.include_router(messages.router, prefix="/v3") app.include_router(conclusions.router, prefix="/v3") app.include_router(keys.router, prefix="/v3") app.include_router(webhooks.router, prefix="/v3") +app.include_router(deriver_metrics.router) # Prometheus metrics endpoint app.add_route("/metrics", metrics_endpoint, methods=["GET"]) diff --git a/src/reconciler/embed_now.py b/src/reconciler/embed_now.py index f76fa760..5308c062 100644 --- a/src/reconciler/embed_now.py +++ b/src/reconciler/embed_now.py @@ -33,7 +33,7 @@ from src.dependencies import tracked_db from src.embedding_client import embedding_client from src.exceptions import VectorStoreError from src.reconciler.sync_vectors import ( - _backoff_eligible, # pyright: ignore[reportPrivateUsage] + backoff_eligible, build_message_vector_record, compute_chunk_positions, ) @@ -177,7 +177,7 @@ async def _claim_and_lease(message_ids: list[str]) -> list[_ClaimedChunk]: and_( models.MessageEmbedding.message_id.in_(message_ids), models.MessageEmbedding.sync_state == "pending", - _backoff_eligible(models.MessageEmbedding.last_sync_at), + backoff_eligible(models.MessageEmbedding.last_sync_at), ) ) .order_by(models.MessageEmbedding.message_id, models.MessageEmbedding.id) diff --git a/src/reconciler/sync_vectors.py b/src/reconciler/sync_vectors.py index 1a8e99b5..b9b06418 100644 --- a/src/reconciler/sync_vectors.py +++ b/src/reconciler/sync_vectors.py @@ -39,7 +39,7 @@ MAX_SYNC_ATTEMPTS = 20 # After this many failures, mark as failed SYNC_BACKOFF = datetime.timedelta(minutes=10) -def _backoff_eligible( +def backoff_eligible( last_sync_at: InstrumentedAttribute[datetime.datetime | None], ) -> ColumnElement[bool]: """Rows are eligible for sync if never attempted or past the backoff window.""" @@ -92,7 +92,7 @@ async def _get_documents_needing_sync( and_( models.Document.deleted_at.is_(None), models.Document.sync_state == "pending", # Only pending items - _backoff_eligible(models.Document.last_sync_at), + backoff_eligible(models.Document.last_sync_at), ) ) .order_by(models.Document.last_sync_at.asc().nullsfirst()) @@ -132,7 +132,7 @@ async def _get_message_embeddings_needing_sync( .where( and_( models.MessageEmbedding.sync_state == "pending", - _backoff_eligible(models.MessageEmbedding.last_sync_at), + backoff_eligible(models.MessageEmbedding.last_sync_at), ) ) .group_by(models.MessageEmbedding.message_id) @@ -153,7 +153,7 @@ async def _get_message_embeddings_needing_sync( and_( models.MessageEmbedding.message_id.in_(message_ids), models.MessageEmbedding.sync_state == "pending", - _backoff_eligible(models.MessageEmbedding.last_sync_at), + backoff_eligible(models.MessageEmbedding.last_sync_at), ) ) .order_by(models.MessageEmbedding.message_id, models.MessageEmbedding.id) diff --git a/src/routers/deriver_metrics.py b/src/routers/deriver_metrics.py new file mode 100644 index 00000000..547a5153 --- /dev/null +++ b/src/routers/deriver_metrics.py @@ -0,0 +1,41 @@ +"""Deriver work metrics as JSON, with the age of the measurement alongside them.""" + +from logging import getLogger + +from fastapi import APIRouter, HTTPException + +from src.backlog import DeriverMetricsPoller + +logger = getLogger(__name__) + +router = APIRouter(prefix="/deriver", tags=["deriver"]) + +_poller: DeriverMetricsPoller | None = None + + +def set_deriver_metrics_poller(poller: DeriverMetricsPoller | None) -> None: + global _poller + _poller = poller + + +@router.get("/metrics") +async def get_deriver_metrics_response() -> dict[str, float | int]: + """Seconds of outstanding deriver work, plus the raw counts behind it.""" + snapshot = _poller.snapshot if _poller is not None else None + if snapshot is None or snapshot.measured_at is None: + raise HTTPException( + status_code=503, detail="No deriver measurement available yet" + ) + + return { + "outstanding_work_seconds": snapshot.signal_seconds, + "eligible_work_units": snapshot.stats.eligible_work_units, + "claimed_work_units": snapshot.stats.claimed_work_units, + "pending_items": snapshot.stats.pending_items, + "oldest_pending_age_seconds": snapshot.stats.oldest_pending_age_seconds, + "embeddings_pending": snapshot.stats.embeddings_pending, + "embeddings_pending_due": snapshot.stats.embeddings_pending_due, + "dreams_due": snapshot.dreams_due, + "measured_at": snapshot.measured_at, + "measurement_age_seconds": snapshot.age_seconds or 0.0, + } diff --git a/src/schemas/__init__.py b/src/schemas/__init__.py index 9f414583..0f93278a 100644 --- a/src/schemas/__init__.py +++ b/src/schemas/__init__.py @@ -78,6 +78,7 @@ from src.schemas.configuration import ( WorkspaceConfiguration, ) from src.schemas.internal import ( + DeriverMetrics, DocumentBase, DocumentCreate, DocumentMetadata, @@ -163,6 +164,7 @@ __all__ = [ "WorkspaceMessageSearchOptions", "WorkspaceUpdate", # internal + "DeriverMetrics", "DocumentBase", "DocumentCreate", "DocumentMetadata", diff --git a/src/schemas/internal.py b/src/schemas/internal.py index 2d299feb..f6399435 100644 --- a/src/schemas/internal.py +++ b/src/schemas/internal.py @@ -140,6 +140,17 @@ class QueueCounts(BaseModel): sessions: dict[str, SessionCounts] +class DeriverMetrics(BaseModel): + """Database-wide view of the deriver's outstanding work.""" + + eligible_work_units: int = 0 + claimed_work_units: int = 0 + pending_items: int = 0 + oldest_pending_age_seconds: float = 0.0 + embeddings_pending: int = 0 + embeddings_pending_due: int = 0 + + class QueueStatusRow(BaseModel): """Represents a row from the queue status SQL query result.""" diff --git a/src/telemetry/prometheus/metrics.py b/src/telemetry/prometheus/metrics.py index cead893a..6859c7cd 100644 --- a/src/telemetry/prometheus/metrics.py +++ b/src/telemetry/prometheus/metrics.py @@ -199,6 +199,69 @@ message_embeddings_pending_gauge = NamespacedGauge( ["namespace"], ) +message_embeddings_pending_due_gauge = NamespacedGauge( + "message_embeddings_pending_due", + "Pending MessageEmbedding rows past their retry backoff, so a sync attempt " + + "is due. Service-wide DB count, reported independently by every API " + + "replica — aggregate with max() or avg(), never sum()", + ["namespace"], +) + +deriver_outstanding_work_seconds_gauge = NamespacedGauge( + "deriver_outstanding_work_seconds", + "Seconds of outstanding deriver work, 0 when a deriver has nothing to do. " + + "Service-wide DB value, reported independently by every API replica — " + + "aggregate with max(), never sum()", + ["namespace"], +) + +deriver_queue_work_units_eligible_gauge = NamespacedGauge( + "deriver_queue_work_units_eligible", + "Work units a deriver could claim right now, ignoring stale claims. " + + "Service-wide DB count, reported independently by every API replica — " + + "aggregate with max() or avg(), never sum()", + ["namespace"], +) + +deriver_queue_work_units_claimed_gauge = NamespacedGauge( + "deriver_queue_work_units_claimed", + "Work units held by a claim refreshed inside the stale timeout, so work is " + + "in flight. Service-wide DB count, reported independently by every API " + + "replica — aggregate with max() or avg(), never sum()", + ["namespace"], +) + +deriver_queue_items_pending_gauge = NamespacedGauge( + "deriver_queue_items_pending", + "Unprocessed queue rows, whether or not they are claimable yet. " + + "Service-wide DB count, reported independently by every API replica — " + + "aggregate with max() or avg(), never sum()", + ["namespace"], +) + +deriver_queue_oldest_pending_age_seconds_gauge = NamespacedGauge( + "deriver_queue_oldest_pending_age_seconds", + "Age of the oldest unprocessed queue row, 0 when the queue is empty. " + + "Service-wide DB value, reported independently by every API replica — " + + "aggregate with max() or avg(), never sum()", + ["namespace"], +) + +dreams_due_gauge = NamespacedGauge( + "dreams_due", + "Collections whose next dream is due and would actually run. " + + "Service-wide DB count, reported independently by every API replica — " + + "aggregate with max() or avg(), never sum()", + ["namespace"], +) + +deriver_metrics_last_success_timestamp_gauge = NamespacedGauge( + "deriver_metrics_last_success_timestamp_seconds", + "Unix time of the last successful deriver-metrics refresh in this replica. " + + "Alert on time() minus this value; a frozen value means the poller stopped", + ["namespace"], +) + # DB connection-pool health. The in-flight gauge counts statements actually # executing on the wire, so checked_out minus in_flight reveals connections held # but parked (the "idle in transaction during an external call" antipattern). @@ -508,6 +571,10 @@ class PrometheusMetrics: self._touch(embed_now_tasks_shed_counter) self.set_embed_now_tasks_in_flight(0) + self.set_deriver_metrics() + self.set_deriver_outstanding_work(seconds=0) + self.set_dreams_due(count=0) + elif instance_type == "deriver": # deriver tokens: only the valid (token_type, component) tuples per # task_type (see _DERIVER_TOKEN_COMBOS_BY_TASK). @@ -548,6 +615,46 @@ class PrometheusMetrics: except Exception as e: self._handle_metric_error("set_message_embeddings_pending", e) + def set_deriver_metrics( + self, + *, + eligible_work_units: int = 0, + claimed_work_units: int = 0, + pending_items: int = 0, + oldest_pending_age_seconds: float = 0.0, + embeddings_pending: int = 0, + embeddings_pending_due: int = 0, + ) -> None: + try: + deriver_queue_work_units_eligible_gauge.labels().set(eligible_work_units) + deriver_queue_work_units_claimed_gauge.labels().set(claimed_work_units) + deriver_queue_items_pending_gauge.labels().set(pending_items) + deriver_queue_oldest_pending_age_seconds_gauge.labels().set( + oldest_pending_age_seconds + ) + message_embeddings_pending_gauge.labels().set(embeddings_pending) + message_embeddings_pending_due_gauge.labels().set(embeddings_pending_due) + except Exception as e: + self._handle_metric_error("set_deriver_metrics", e) + + def set_deriver_outstanding_work(self, *, seconds: float) -> None: + try: + deriver_outstanding_work_seconds_gauge.labels().set(seconds) + except Exception as e: + self._handle_metric_error("set_deriver_outstanding_work", e) + + def set_dreams_due(self, *, count: int) -> None: + try: + dreams_due_gauge.labels().set(count) + except Exception as e: + self._handle_metric_error("set_dreams_due", e) + + def set_deriver_metrics_last_success(self, *, timestamp: float) -> None: + try: + deriver_metrics_last_success_timestamp_gauge.labels().set(timestamp) + except Exception as e: + self._handle_metric_error("set_deriver_metrics_last_success", e) + prometheus_metrics = PrometheusMetrics() diff --git a/tests/crud/test_deriver_metrics_query.py b/tests/crud/test_deriver_metrics_query.py new file mode 100644 index 00000000..a67c17a8 --- /dev/null +++ b/tests/crud/test_deriver_metrics_query.py @@ -0,0 +1,394 @@ +import datetime + +import pytest +from nanoid import generate as generate_nanoid +from sqlalchemy.ext.asyncio import AsyncSession + +from src import crud, models +from src.config import settings + +pytestmark = pytest.mark.asyncio + + +async def _make_session( + db: AsyncSession, workspace: models.Workspace +) -> models.Session: + session = models.Session(name=str(generate_nanoid()), workspace_name=workspace.name) + db.add(session) + await db.flush() + return session + + +async def _add_representation_item( + db: AsyncSession, + workspace: models.Workspace, + peer: models.Peer, + session: models.Session, + *, + work_unit_key: str, + token_count: int, + age_seconds: int = 0, + seq: int = 1, +) -> models.QueueItem: + message = models.Message( + session_name=session.name, + content="x", + token_count=token_count, + seq_in_session=seq, + peer_name=peer.name, + workspace_name=workspace.name, + ) + db.add(message) + await db.flush() + + item = models.QueueItem( + session_id=session.id, + work_unit_key=work_unit_key, + task_type="representation", + payload={}, + processed=False, + workspace_name=workspace.name, + message_id=message.id, + created_at=datetime.datetime.now(datetime.UTC) + - datetime.timedelta(seconds=age_seconds), + ) + db.add(item) + await db.flush() + return item + + +async def _add_message( + db: AsyncSession, + workspace: models.Workspace, + peer: models.Peer, + session: models.Session, + *, + seq: int = 1, +) -> models.Message: + message = models.Message( + session_name=session.name, + content="x", + token_count=1, + seq_in_session=seq, + peer_name=peer.name, + workspace_name=workspace.name, + ) + db.add(message) + await db.flush() + return message + + +def _stale_timestamp() -> datetime.datetime: + return datetime.datetime.now(datetime.UTC) - datetime.timedelta( + minutes=settings.DERIVER.STALE_SESSION_TIMEOUT_MINUTES + 1 + ) + + +class TestDeriverMetrics: + async def test_empty_queue_reports_zero( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], # pyright: ignore[reportUnusedParameter] + ): + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 0 + assert stats.claimed_work_units == 0 + assert stats.pending_items == 0 + assert stats.oldest_pending_age_seconds == 0.0 + + async def test_sub_threshold_batch_is_pending_but_not_eligible( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A small, fresh batch is real work that a deriver would not yet claim.""" + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:small", + token_count=1, + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.pending_items == 1 + assert stats.eligible_work_units == 0 + + async def test_token_threshold_makes_batch_eligible( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:big", + token_count=settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS, + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 1 + + async def test_age_flush_makes_sub_threshold_batch_eligible( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:old", + token_count=1, + age_seconds=settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS + 60, + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 1 + assert stats.oldest_pending_age_seconds >= ( + settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS + ) + + async def test_non_representation_work_is_eligible_immediately( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, _peer = sample_data + + db_session.add( + models.QueueItem( + work_unit_key="reconciler:sync_vectors", + task_type="reconciler", + payload={}, + processed=False, + workspace_name=workspace.name, + ) + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 1 + + async def test_live_claim_is_counted_as_work_in_flight( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A claimed work unit is not claimable, but it is still outstanding work.""" + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:claimed", + token_count=settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS, + ) + db_session.add( + models.ActiveQueueSession(work_unit_key="representation:claimed") + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 0 + assert stats.claimed_work_units == 1 + + async def test_stale_claim_does_not_hide_work_and_is_not_in_flight( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A dead worker's claim must not read as in flight, and must not hide work.""" + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:abandoned", + token_count=settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS, + ) + db_session.add( + models.ActiveQueueSession( + work_unit_key="representation:abandoned", + last_updated=_stale_timestamp(), + ) + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.eligible_work_units == 1 + assert stats.claimed_work_units == 0 + + async def test_processed_items_are_not_counted( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + item = await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:done", + token_count=settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS, + ) + item.processed = True + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.pending_items == 0 + assert stats.eligible_work_units == 0 + assert stats.oldest_pending_age_seconds == 0.0 + + +class TestPendingEmbeddings: + async def test_never_attempted_row_is_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + message = await _add_message(db_session, workspace, peer, session) + db_session.add( + models.MessageEmbedding( + content="x", + message_id=message.public_id, + workspace_name=workspace.name, + session_name=session.name, + peer_name=peer.name, + sync_state="pending", + ) + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.embeddings_pending == 1 + assert stats.embeddings_pending_due == 1 + + async def test_row_inside_its_retry_wait_is_pending_but_not_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A backing-off row is work the deriver cannot act on yet.""" + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + message = await _add_message(db_session, workspace, peer, session) + db_session.add( + models.MessageEmbedding( + content="x", + message_id=message.public_id, + workspace_name=workspace.name, + session_name=session.name, + peer_name=peer.name, + sync_state="pending", + last_sync_at=datetime.datetime.now(datetime.UTC), + sync_attempts=1, + ) + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.embeddings_pending == 1 + assert stats.embeddings_pending_due == 0 + + async def test_synced_rows_are_not_counted( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + message = await _add_message(db_session, workspace, peer, session) + db_session.add( + models.MessageEmbedding( + content="x", + message_id=message.public_id, + workspace_name=workspace.name, + session_name=session.name, + peer_name=peer.name, + sync_state="synced", + ) + ) + await db_session.commit() + + stats = await crud.get_deriver_metrics(db_session) + + assert stats.embeddings_pending == 0 + assert stats.embeddings_pending_due == 0 + + +class TestMetricsAgreeWithDeriver: + @pytest.mark.parametrize( + "token_count,age_seconds", + [ + (1, 0), + (settings.DERIVER.REPRESENTATION_BATCH_WORK_UNIT_TARGET_TOKENS, 0), + (1, settings.DERIVER.REPRESENTATION_BATCH_MAX_AGE_SECONDS + 60), + ], + ids=["sub-threshold", "token-threshold", "age-flush"], + ) + async def test_eligible_count_matches_what_the_deriver_claims( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + token_count: int, + age_seconds: int, + ): + """The gauge is only trustworthy if it uses the deriver's own rule.""" + from src.deriver.queue_manager import QueueManager + + workspace, peer = sample_data + session = await _make_session(db_session, workspace) + + await _add_representation_item( + db_session, + workspace, + peer, + session, + work_unit_key="representation:agreement", + token_count=token_count, + age_seconds=age_seconds, + ) + await db_session.commit() + + expected = (await crud.get_deriver_metrics(db_session)).eligible_work_units + claimed = await QueueManager().get_and_claim_work_units() + + assert len(claimed) == expected diff --git a/tests/dreamer/test_dream_due.py b/tests/dreamer/test_dream_due.py new file mode 100644 index 00000000..dad75de6 --- /dev/null +++ b/tests/dreamer/test_dream_due.py @@ -0,0 +1,321 @@ +"""Tests for the read-only count of collections whose next dream is due.""" + +import datetime +from unittest.mock import patch + +import pytest +from nanoid import generate as generate_nanoid +from sqlalchemy.ext.asyncio import AsyncSession + +from src import models +from src.dreamer.dream_due import count_due_dreams +from src.schemas import DreamType +from src.utils.work_unit import construct_work_unit_key + + +def _now() -> datetime.datetime: + return datetime.datetime.now(datetime.UTC) + + +async def _make_collection( + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + internal_metadata: dict[str, object] | None = None, +) -> models.Collection: + workspace, peer = sample_data + collection = models.Collection( + observer=peer.name, + observed=peer.name, + workspace_name=workspace.name, + internal_metadata=internal_metadata or {}, + ) + db_session.add(collection) + await db_session.commit() + return collection + + +async def _make_session( + db_session: AsyncSession, + workspace_name: str, + configuration: dict[str, object] | None = None, +) -> str: + session = models.Session( + name=f"s-{generate_nanoid()}", + workspace_name=workspace_name, + configuration=configuration or {}, + ) + db_session.add(session) + await db_session.commit() + return session.name + + +async def _insert_docs( + db_session: AsyncSession, + collection: models.Collection, + level: str, + count: int, + *, + age_minutes: int = 0, + session_name: str | None = None, + sessionless: bool = False, +) -> None: + if session_name is None and not sessionless: + session_name = await _make_session(db_session, collection.workspace_name) + created_at = _now() - datetime.timedelta(minutes=age_minutes) + for _ in range(count): + db_session.add( + models.Document( + content="test", + level=level, + workspace_name=collection.workspace_name, + observer=collection.observer, + observed=collection.observed, + session_name=session_name, + created_at=created_at, + ) + ) + await db_session.commit() + + +async def _insert_dream_item( + db_session: AsyncSession, + collection: models.Collection, + *, + age_minutes: int, + processed: bool, + error: str | None = None, +) -> None: + work_unit_key = construct_work_unit_key( + collection.workspace_name, + { + "task_type": "dream", + "observer": collection.observer, + "observed": collection.observed, + "dream_type": DreamType.OMNI.value, + }, + ) + db_session.add( + models.QueueItem( + work_unit_key=work_unit_key, + payload={"task_type": "dream"}, + task_type="dream", + workspace_name=collection.workspace_name, + processed=processed, + error=error, + created_at=_now() - datetime.timedelta(minutes=age_minutes), + ) + ) + await db_session.commit() + + +@pytest.fixture(autouse=True) +def _pin_dream_config(): # pyright: ignore[reportUnusedFunction] + with ( + patch("src.dreamer.dream_due.settings.DREAM.ENABLED", True), + patch("src.dreamer.dream_due.settings.DREAM.DOCUMENT_THRESHOLD", 50), + patch("src.dreamer.dream_due.settings.DREAM.ENABLED_TYPES", ["omni"]), + patch("src.dreamer.dream_due.settings.DREAM.IDLE_TIMEOUT_MINUTES", 60), + patch("src.dreamer.dream_due.settings.DREAM.MIN_HOURS_BETWEEN_DREAMS", 8), + ): + yield + + +@pytest.mark.asyncio +class TestCountDueDreams: + async def test_below_threshold_is_not_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 30, age_minutes=90) + + assert await count_due_dreams(db_session) == 0 + + async def test_derived_levels_do_not_count( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 30, age_minutes=90) + await _insert_docs(db_session, collection, "deductive", 40, age_minutes=90) + + assert await count_due_dreams(db_session) == 0 + + async def test_threshold_met_but_not_idle_is_not_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A collection still receiving documents is not idle yet.""" + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=1) + + assert await count_due_dreams(db_session) == 0 + + async def test_threshold_met_and_idle_is_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + assert await count_due_dreams(db_session) == 1 + + async def test_documents_since_last_dream_uses_stored_count( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection( + db_session, sample_data, {"dream": {"last_dream_document_count": 40}} + ) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + assert await count_due_dreams(db_session) == 0 + + async def test_min_hours_gate_blocks_a_recent_dream( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + last_dream_at = (_now() - datetime.timedelta(hours=2)).isoformat() + collection = await _make_collection( + db_session, sample_data, {"dream": {"last_dream_at": last_dream_at}} + ) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + assert await count_due_dreams(db_session) == 0 + + async def test_naive_last_dream_at_is_read_as_utc( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A stored timestamp with no offset must gate, not raise.""" + naive = (_now() - datetime.timedelta(hours=2)).replace(tzinfo=None).isoformat() + collection = await _make_collection( + db_session, sample_data, {"dream": {"last_dream_at": naive}} + ) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + assert await count_due_dreams(db_session) == 0 + + async def test_pending_dream_item_blocks( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + await _insert_dream_item( + db_session, collection, age_minutes=10, processed=False + ) + + assert await count_due_dreams(db_session) == 0 + + async def test_failed_dream_waits_for_new_documents( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """Without this the count never returns to zero.""" + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + await _insert_dream_item( + db_session, collection, age_minutes=80, processed=True, error="boom" + ) + + assert await count_due_dreams(db_session) == 0 + + async def test_failed_dream_retries_after_new_documents( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + await _insert_dream_item( + db_session, collection, age_minutes=80, processed=True, error="boom" + ) + await _insert_docs(db_session, collection, "explicit", 1, age_minutes=70) + + assert await count_due_dreams(db_session) == 1 + + async def test_sessionless_documents_are_not_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """The deriver's own enqueue path refuses these, so they must not count.""" + collection = await _make_collection(db_session, sample_data) + await _insert_docs( + db_session, collection, "explicit", 60, age_minutes=90, sessionless=True + ) + + assert await count_due_dreams(db_session) == 0 + + async def test_newest_document_decides_the_session( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=120) + + assert await count_due_dreams(db_session) == 1 + + await _insert_docs( + db_session, collection, "explicit", 1, age_minutes=90, sessionless=True + ) + + assert await count_due_dreams(db_session) == 0 + + async def test_session_with_dreams_disabled_is_not_due( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + """A dream the enqueue path would refuse must not be counted.""" + collection = await _make_collection(db_session, sample_data) + session_name = await _make_session( + db_session, + collection.workspace_name, + {"dream": {"enabled": False}}, + ) + await _insert_docs( + db_session, + collection, + "explicit", + 60, + age_minutes=90, + session_name=session_name, + ) + + assert await count_due_dreams(db_session) == 0 + + async def test_dreams_disabled_globally_returns_zero( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + with patch("src.dreamer.dream_due.settings.DREAM.ENABLED", False): + assert await count_due_dreams(db_session) == 0 + + async def test_card_refresh_is_never_counted( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + ): + collection = await _make_collection(db_session, sample_data) + await _insert_docs(db_session, collection, "explicit", 60, age_minutes=90) + + with patch( + "src.dreamer.dream_due.settings.DREAM.ENABLED_TYPES", ["card_refresh"] + ): + assert await count_due_dreams(db_session) == 0 diff --git a/tests/telemetry/test_metric_zero_init.py b/tests/telemetry/test_metric_zero_init.py index e69f50df..eb412287 100644 --- a/tests/telemetry/test_metric_zero_init.py +++ b/tests/telemetry/test_metric_zero_init.py @@ -131,6 +131,19 @@ def test_deriver_token_combos_are_valid_and_complete(): ) not in ingestion +_API_DERIVER_METRIC_GAUGES = ( + "deriver_outstanding_work_seconds", + "deriver_queue_work_units_eligible", + "deriver_queue_work_units_claimed", + "deriver_queue_items_pending", + "deriver_queue_oldest_pending_age_seconds", + "dreams_due", + "message_embeddings_pending_due", +) + +_SHARED_DERIVER_METRIC_GAUGES = ("message_embeddings_pending",) + + # --------------------------------------------------------------------------- # API-process zero-init # --------------------------------------------------------------------------- @@ -161,6 +174,8 @@ def test_api_init_materializes_dialectic_and_embed(): ) assert sample("embed_now_tasks_shed_total") is not None assert sample("embed_now_tasks_in_flight") == 0.0 # gauge, explicit .set(0) + for gauge in (*_API_DERIVER_METRIC_GAUGES, *_SHARED_DERIVER_METRIC_GAUGES): + assert sample(gauge) == 0.0, f"{gauge} was not zero-initialized" @pytest.mark.usefixtures("metrics_enabled") @@ -310,6 +325,9 @@ def test_deriver_init_does_not_touch_api_counters(): # the API-process embed_now counters are equally off-limits assert sample("embed_now_tasks_shed_total") is None assert sample("embed_now_tasks_in_flight") is None + # so are the deriver-work gauges: the deriver never measures its own backlog + for gauge in _API_DERIVER_METRIC_GAUGES: + assert sample(gauge) is None, f"{gauge} must be API-only" # --------------------------------------------------------------------------- diff --git a/tests/test_deriver_metrics.py b/tests/test_deriver_metrics.py new file mode 100644 index 00000000..acd97330 --- /dev/null +++ b/tests/test_deriver_metrics.py @@ -0,0 +1,207 @@ +"""Tests for the outstanding-work value, the poller and the JSON route.""" + +import time +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import HTTPException + +from src import schemas +from src.backlog import ( + DeriverMetricsPoller, + DeriverMetricsSnapshot, + active_work_seconds, + outstanding_work_seconds, +) +from src.routers import deriver_metrics + + +class TestScaleSignal: + def test_nothing_outstanding_reads_zero(self): + assert outstanding_work_seconds(schemas.DeriverMetrics(), dreams_due=0) == 0.0 + + def test_claimable_work_reports_the_active_value(self): + stats = schemas.DeriverMetrics(eligible_work_units=1) + + assert outstanding_work_seconds(stats, dreams_due=0) == active_work_seconds() + + def test_work_in_flight_still_reports_the_active_value(self): + """A row claimed a moment ago has a small age and would read as idle.""" + stats = schemas.DeriverMetrics( + claimed_work_units=1, pending_items=1, oldest_pending_age_seconds=2.0 + ) + + assert outstanding_work_seconds(stats, dreams_due=0) == active_work_seconds() + + def test_waiting_batch_reports_its_real_age(self): + """The real age is what tells a caller how close the flush is.""" + stats = schemas.DeriverMetrics( + pending_items=3, oldest_pending_age_seconds=1234.0 + ) + + assert outstanding_work_seconds(stats, dreams_due=0) == 1234.0 + + def test_embeddings_due_an_attempt_report_the_active_value(self): + stats = schemas.DeriverMetrics(embeddings_pending=5, embeddings_pending_due=5) + + assert outstanding_work_seconds(stats, dreams_due=0) == active_work_seconds() + + def test_embeddings_inside_their_retry_wait_do_not(self): + """Otherwise one permanently failing row holds the value up for hours.""" + stats = schemas.DeriverMetrics(embeddings_pending=5) + + assert outstanding_work_seconds(stats, dreams_due=0) == 0.0 + + def test_a_due_dream_reports_the_active_value(self): + assert ( + outstanding_work_seconds(schemas.DeriverMetrics(), dreams_due=1) + == active_work_seconds() + ) + + def test_active_value_is_positive(self): + assert active_work_seconds() > 0 + + +@pytest.mark.asyncio +class TestPoller: + async def test_refresh_publishes_a_snapshot(self): + stats = schemas.DeriverMetrics(eligible_work_units=2, pending_items=4) + poller = DeriverMetricsPoller() + + with ( + patch( + "src.backlog.crud.get_deriver_metrics", + AsyncMock(return_value=stats), + ), + patch("src.backlog.count_due_dreams", AsyncMock(return_value=3)), + ): + await poller.refresh() + + snapshot = poller.snapshot + assert snapshot.measured_at is not None + assert snapshot.stats.eligible_work_units == 2 + assert snapshot.dreams_due == 3 + assert snapshot.signal_seconds == active_work_seconds() + + async def test_dream_query_runs_on_its_own_spacing(self): + """The dream query is the expensive one, so it must not run every pass.""" + stats = schemas.DeriverMetrics() + poller = DeriverMetricsPoller() + dream_count = AsyncMock(return_value=1) + + with ( + patch( + "src.backlog.crud.get_deriver_metrics", + AsyncMock(return_value=stats), + ), + patch("src.backlog.count_due_dreams", dream_count), + ): + await poller.refresh() + await poller.refresh() + + assert dream_count.await_count == 1 + assert poller.snapshot.dreams_due == 1 + + async def test_a_failed_dream_query_is_retried_on_the_next_pass(self): + """Advancing the deadline first would republish the old count for a whole interval.""" + stats = schemas.DeriverMetrics() + poller = DeriverMetricsPoller() + dream_count = AsyncMock(side_effect=[RuntimeError("db down"), 4]) + + with ( + patch( + "src.backlog.crud.get_deriver_metrics", + AsyncMock(return_value=stats), + ), + patch("src.backlog.count_due_dreams", dream_count), + ): + with pytest.raises(RuntimeError): + await poller.refresh() + await poller.refresh() + + assert dream_count.await_count == 2 + assert poller.snapshot.dreams_due == 4 + + async def test_a_failed_pass_leaves_the_previous_snapshot_alone(self): + """A half-finished pass must never be published as a measurement.""" + stats = schemas.DeriverMetrics(eligible_work_units=1) + poller = DeriverMetricsPoller() + + with ( + patch( + "src.backlog.crud.get_deriver_metrics", + AsyncMock(return_value=stats), + ), + patch("src.backlog.count_due_dreams", AsyncMock(return_value=0)), + ): + await poller.refresh() + + first = poller.snapshot + + with ( + patch( + "src.backlog.crud.get_deriver_metrics", + AsyncMock(side_effect=RuntimeError("db down")), + ), + pytest.raises(RuntimeError), + ): + await poller.refresh() + + assert poller.snapshot is first + + +@pytest.mark.asyncio +class TestDeriverMetricsRoute: + async def test_serves_the_cached_snapshot(self): + poller = DeriverMetricsPoller() + poller._snapshot = DeriverMetricsSnapshot( # pyright: ignore[reportPrivateUsage] + signal_seconds=1800.0, + dreams_due=1, + stats=schemas.DeriverMetrics(eligible_work_units=2, pending_items=5), + measured_at=time.time(), + ) + deriver_metrics.set_deriver_metrics_poller(poller) + try: + body = await deriver_metrics.get_deriver_metrics_response() + finally: + deriver_metrics.set_deriver_metrics_poller(None) + + assert body["outstanding_work_seconds"] == 1800.0 + assert body["eligible_work_units"] == 2 + assert body["pending_items"] == 5 + assert body["dreams_due"] == 1 + + async def test_errors_before_the_first_pass(self): + """A 503 tells the caller there is no measurement; a 0 would be a lie.""" + deriver_metrics.set_deriver_metrics_poller(DeriverMetricsPoller()) + try: + with pytest.raises(HTTPException) as excinfo: + await deriver_metrics.get_deriver_metrics_response() + finally: + deriver_metrics.set_deriver_metrics_poller(None) + + assert excinfo.value.status_code == 503 + + async def test_serves_an_old_snapshot_with_its_age(self): + """The caller decides what is too old, from measurement_age_seconds.""" + poller = DeriverMetricsPoller() + poller._snapshot = DeriverMetricsSnapshot( # pyright: ignore[reportPrivateUsage] + signal_seconds=7.0, + measured_at=time.time() - 3600, + ) + deriver_metrics.set_deriver_metrics_poller(poller) + try: + body = await deriver_metrics.get_deriver_metrics_response() + finally: + deriver_metrics.set_deriver_metrics_poller(None) + + assert body["outstanding_work_seconds"] == 7.0 + assert body["measurement_age_seconds"] >= 3600 + + async def test_errors_when_no_poller_is_registered(self): + deriver_metrics.set_deriver_metrics_poller(None) + + with pytest.raises(HTTPException) as excinfo: + await deriver_metrics.get_deriver_metrics_response() + + assert excinfo.value.status_code == 503