diff --git a/src/config.py b/src/config.py index 0e39fdcc..e8b8461e 100644 --- a/src/config.py +++ b/src/config.py @@ -754,6 +754,10 @@ class EmbeddingSettings(HonchoSettings): # Caps concurrent message-embedding fan-out on the API request path (the # immediate-embed background task). The reconciler is unaffected. MAX_CONCURRENT_EMBEDDINGS: Annotated[int, Field(default=10, gt=0, le=100)] = 10 + # Caps in-flight immediate-embed background tasks per API process. When + # saturated, message creation skips the fast path entirely and the + # reconciler embeds on its next cycle. 0 disables the fast path. + MAX_PENDING_EMBED_TASKS: Annotated[int, Field(default=50, ge=0)] = 50 @model_validator(mode="before") @classmethod diff --git a/src/reconciler/embed_now.py b/src/reconciler/embed_now.py index 12c71065..f958d3b9 100644 --- a/src/reconciler/embed_now.py +++ b/src/reconciler/embed_now.py @@ -23,6 +23,7 @@ import logging from dataclasses import dataclass from typing import Any +from fastapi import BackgroundTasks from sqlalchemy import and_, func, select, update from sqlalchemy.ext.asyncio import AsyncSession @@ -36,6 +37,7 @@ from src.reconciler.sync_vectors import ( build_message_vector_record, compute_chunk_positions, ) +from src.telemetry import prometheus_metrics from src.telemetry.events import EmbeddingCallPurpose from src.utils.types import embedding_call_purpose from src.vector_store import VectorRecord, VectorStore, get_external_vector_store @@ -66,6 +68,46 @@ def reset_embed_semaphore() -> None: _embed_semaphore = None +class EmbedTaskGate: + """Non-blocking admission gate for immediate-embed background tasks. + + Bounds the number of in-flight tasks per API process at + ``EMBEDDING.MAX_PENDING_EMBED_TASKS``. When saturated, nothing is scheduled: + the rows are already ``sync_state='pending'``, so the reconciler embeds them + on its next cycle. The count is taken at schedule time (not task start) + because background tasks only run after the response is sent — a request + burst would otherwise stack up unbounded scheduled-but-not-started tasks. + + ``in_flight`` is mutated only from the event loop (request handlers and the + tracked task), so a plain int is race-free. + """ + + def __init__(self) -> None: + self.in_flight: int = 0 + + def try_schedule( + self, background_tasks: BackgroundTasks, message_ids: list[str] + ) -> bool: + """Schedule ``embed_messages_now`` if the cap allows it; return whether + the task was scheduled.""" + if self.in_flight >= settings.EMBEDDING.MAX_PENDING_EMBED_TASKS: + if settings.METRICS.ENABLED: + prometheus_metrics.record_embed_now_task_shed() + return False + self.in_flight += 1 + background_tasks.add_task(self._run, message_ids) + return True + + async def _run(self, message_ids: list[str]) -> None: + try: + await embed_messages_now(message_ids) + finally: + self.in_flight -= 1 + + +embed_task_gate = EmbedTaskGate() + + @dataclass(frozen=True) class _ClaimedChunk: """Plain snapshot of a claimed ``MessageEmbedding`` row. diff --git a/src/routers/messages.py b/src/routers/messages.py index 9ace70cc..b0dd8752 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -21,7 +21,7 @@ from src.config import settings from src.dependencies import db, read_db from src.deriver import enqueue from src.exceptions import FileTooLargeError, ResourceNotFoundException -from src.reconciler.embed_now import embed_messages_now +from src.reconciler.embed_now import embed_task_gate from src.security import require_auth from src.telemetry import prometheus_metrics from src.telemetry.events import FileUploadedEvent, MessageCreatedEvent, emit @@ -161,11 +161,17 @@ async def create_messages_for_session( background_tasks.add_task(enqueue, payloads) # Embed immediately so messages are searchable within seconds; the - # reconciler is the fallback for anything left pending. + # reconciler is the fallback for anything left pending. Scheduling is + # capped per process — when saturated, the reconciler picks them up. if settings.EMBED_MESSAGES and created_messages: - background_tasks.add_task( - embed_messages_now, [m.public_id for m in created_messages] + scheduled = embed_task_gate.try_schedule( + background_tasks, [m.public_id for m in created_messages] ) + if not scheduled: + logger.debug( + "Immediate-embed tasks saturated; deferring %s message(s) to reconciler", + len(created_messages), + ) return created_messages except ValueError as e: @@ -240,11 +246,17 @@ async def create_messages_with_file( background_tasks.add_task(enqueue, payloads) # Embed immediately so messages are searchable within seconds; the - # reconciler is the fallback for anything left pending. + # reconciler is the fallback for anything left pending. Scheduling is + # capped per process — when saturated, the reconciler picks them up. if settings.EMBED_MESSAGES and created_messages: - background_tasks.add_task( - embed_messages_now, [m.public_id for m in created_messages] + scheduled = embed_task_gate.try_schedule( + background_tasks, [m.public_id for m in created_messages] ) + if not scheduled: + logger.debug( + "Immediate-embed tasks saturated; deferring %s message(s) to reconciler", + len(created_messages), + ) logger.debug( "Batch of %s messages created from file uploads and queued for processing", diff --git a/src/telemetry/prometheus/metrics.py b/src/telemetry/prometheus/metrics.py index 9038b9d0..7317c22a 100644 --- a/src/telemetry/prometheus/metrics.py +++ b/src/telemetry/prometheus/metrics.py @@ -87,6 +87,12 @@ messages_created_counter = NamespacedCounter( ["namespace", "workspace_name"], ) +embed_now_tasks_shed_counter = NamespacedCounter( + "embed_now_tasks_shed", + "Immediate-embed background tasks skipped because MAX_PENDING_EMBED_TASKS was reached", + ["namespace"], +) + dialectic_calls_counter = NamespacedCounter( "dialectic_calls", "Total dialectic calls", @@ -204,6 +210,12 @@ class PrometheusMetrics: except Exception as e: self._handle_metric_error("record_messages_created", e) + def record_embed_now_task_shed(self) -> None: + try: + embed_now_tasks_shed_counter.labels().inc() + except Exception as e: + self._handle_metric_error("record_embed_now_task_shed", e) + def record_dialectic_call( self, *, diff --git a/tests/deriver/test_embed_now.py b/tests/deriver/test_embed_now.py index f6acd1c4..dcc055fe 100644 --- a/tests/deriver/test_embed_now.py +++ b/tests/deriver/test_embed_now.py @@ -9,12 +9,18 @@ creates committed fixture rows and asserts on the result via the provided sessio from unittest.mock import AsyncMock, patch import pytest +from fastapi import BackgroundTasks from nanoid import generate as generate_nanoid from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker from src import models -from src.reconciler.embed_now import embed_messages_now, reset_embed_semaphore +from src.config import settings +from src.reconciler.embed_now import ( + embed_messages_now, + embed_task_gate, + reset_embed_semaphore, +) from src.vector_store import VectorStore @@ -65,10 +71,66 @@ async def _create_message_with_pending_chunks( @pytest.fixture(autouse=True) def reset_semaphore_fixture(): - """Rebuild the module semaphore per test so it binds to the active loop.""" + """Rebuild the module semaphore per test so it binds to the active loop, + and clear the admission gate's in-flight count.""" reset_embed_semaphore() + embed_task_gate.in_flight = 0 yield reset_embed_semaphore() + embed_task_gate.in_flight = 0 + + +@pytest.mark.asyncio +class TestEmbedTaskGate: + """Admission gate for immediate-embed background tasks + (EMBEDDING.MAX_PENDING_EMBED_TASKS).""" + + async def test_admits_under_cap_and_releases_slot(self) -> None: + """Under the cap, the task is scheduled; running it embeds the given ids + and releases the slot.""" + tasks = BackgroundTasks() + with ( + patch.object(settings.EMBEDDING, "MAX_PENDING_EMBED_TASKS", 2), + patch( + "src.reconciler.embed_now.embed_messages_now", new=AsyncMock() + ) as mock_embed, + ): + assert embed_task_gate.try_schedule(tasks, ["msg_1"]) is True + assert embed_task_gate.in_flight == 1 + await tasks() + mock_embed.assert_awaited_once_with(["msg_1"]) + assert embed_task_gate.in_flight == 0 + + async def test_rejects_at_cap_without_scheduling(self) -> None: + """At the cap, nothing is scheduled and False is returned.""" + tasks = BackgroundTasks() + with patch.object(settings.EMBEDDING, "MAX_PENDING_EMBED_TASKS", 1): + assert embed_task_gate.try_schedule(tasks, ["msg_1"]) is True + assert embed_task_gate.try_schedule(tasks, ["msg_2"]) is False + assert len(tasks.tasks) == 1 + assert embed_task_gate.in_flight == 1 + + async def test_slot_released_when_task_raises(self) -> None: + """A failing task still releases its slot (finally path).""" + tasks = BackgroundTasks() + with ( + patch.object(settings.EMBEDDING, "MAX_PENDING_EMBED_TASKS", 1), + patch( + "src.reconciler.embed_now.embed_messages_now", + new=AsyncMock(side_effect=RuntimeError("boom")), + ), + ): + assert embed_task_gate.try_schedule(tasks, ["msg_1"]) is True + with pytest.raises(RuntimeError): + await tasks() + assert embed_task_gate.in_flight == 0 + + async def test_zero_cap_disables_fast_path(self) -> None: + """MAX_PENDING_EMBED_TASKS=0 rejects every schedule attempt.""" + tasks = BackgroundTasks() + with patch.object(settings.EMBEDDING, "MAX_PENDING_EMBED_TASKS", 0): + assert embed_task_gate.try_schedule(tasks, ["msg_1"]) is False + assert len(tasks.tasks) == 0 @pytest.mark.asyncio diff --git a/tests/routes/test_messages.py b/tests/routes/test_messages.py index 3539a735..e70c0f18 100644 --- a/tests/routes/test_messages.py +++ b/tests/routes/test_messages.py @@ -64,7 +64,7 @@ async def test_create_message_schedules_immediate_embed( with ( patch("src.config.settings.EMBED_MESSAGES", True), patch( - "src.routers.messages.embed_messages_now", new=AsyncMock() + "src.reconciler.embed_now.embed_messages_now", new=AsyncMock() ) as mock_embed_now, ): response = client.post( @@ -91,7 +91,35 @@ async def test_create_message_skips_embed_when_disabled( with ( patch("src.config.settings.EMBED_MESSAGES", False), patch( - "src.routers.messages.embed_messages_now", new=AsyncMock() + "src.reconciler.embed_now.embed_messages_now", new=AsyncMock() + ) as mock_embed_now, + ): + response = client.post( + f"/v3/workspaces/{test_workspace.name}/sessions/{test_session.name}/messages", + json={"messages": [{"content": "hello", "peer_id": test_peer.name}]}, + ) + assert response.status_code == 201 + mock_embed_now.assert_not_called() + + +@pytest.mark.asyncio +async def test_create_message_defers_embed_when_saturated( + client: TestClient, db_session: AsyncSession, sample_data: tuple[Workspace, Peer] +): + """When the immediate-embed task cap is saturated, message creation still + succeeds and no embed task runs — rows stay pending for the reconciler.""" + test_workspace, test_peer = sample_data + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + with ( + patch("src.config.settings.EMBED_MESSAGES", True), + patch.object(settings.EMBEDDING, "MAX_PENDING_EMBED_TASKS", 0), + patch( + "src.reconciler.embed_now.embed_messages_now", new=AsyncMock() ) as mock_embed_now, ): response = client.post( @@ -120,7 +148,7 @@ async def test_file_upload_schedules_immediate_embed( with ( patch("src.config.settings.EMBED_MESSAGES", True), patch( - "src.routers.messages.embed_messages_now", new=AsyncMock() + "src.reconciler.embed_now.embed_messages_now", new=AsyncMock() ) as mock_embed_now, ): files = {"file": ("note.txt", io.BytesIO(b"hello world"), "text/plain")}