feat(api): Export deriver backlog as metrics from API endpoint (#1115)

This commit is contained in:
Ulysse Pence 2026-09-02 13:42:48 -04:00 committed by GitHub
parent 997b4764b9
commit 5d992bc65a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
17 changed files with 1656 additions and 38 deletions

146
src/backlog.py Normal file
View File

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

View File

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

View File

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

View File

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

View File

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

216
src/dreamer/dream_due.py Normal file
View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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