diff --git a/src/crud/document.py b/src/crud/document.py index 892925c4..c2fab22c 100644 --- a/src/crud/document.py +++ b/src/crud/document.py @@ -214,6 +214,7 @@ def _uses_pgvector() -> bool: # Shared by is_rejected_duplicate and create_documents candidate resolution. _SEMANTIC_DUP_MAX_DISTANCE = 0.05 _SEMANTIC_DUP_TOP_K = 1 +_SEMANTIC_CANDIDATE_CONCURRENCY = 8 def _semantic_dup_filters(doc: schemas.DocumentCreate) -> dict[str, Any] | None: @@ -538,27 +539,40 @@ async def create_documents( docs_with_embeddings: list[tuple[models.Document, list[float]]] = [] # Resolve external-store dup candidates before the first DB statement. - # None falls back to in-place resolution (pgvector, or no merge partner). + # None = pgvector in-place fallback; [] = skip semantic (no external I/O under db). semantic_candidates: list[list[str] | None] = [None] * len(documents) if deduplicate and not _uses_pgvector(): + resolve_sem = asyncio.Semaphore(_SEMANTIC_CANDIDATE_CONCURRENCY) async def _resolve_candidates(index: int, doc: schemas.DocumentCreate) -> None: filters = _semantic_dup_filters(doc) if filters is None or not doc.embedding: + semantic_candidates[index] = [] return - semantic_candidates[index] = await query_external_vector_document_ids( - workspace_name=workspace_name, - observer=observer, - observed=observed, - embedding=doc.embedding, - top_k=_SEMANTIC_DUP_TOP_K, - max_distance=_SEMANTIC_DUP_MAX_DISTANCE, - filters=filters, - ) + async with resolve_sem: + try: + ids = await query_external_vector_document_ids( + workspace_name=workspace_name, + observer=observer, + observed=observed, + embedding=doc.embedding, + top_k=_SEMANTIC_DUP_TOP_K, + max_distance=_SEMANTIC_DUP_MAX_DISTANCE, + filters=filters, + ) + except Exception: + logger.exception( + "External semantic-candidate resolve failed for %s/%s/%s", + workspace_name, + observer, + observed, + ) + semantic_candidates[index] = [] + return + semantic_candidates[index] = ids or [] await asyncio.gather( - *(_resolve_candidates(i, doc) for i, doc in enumerate(documents)), - return_exceptions=True, + *(_resolve_candidates(i, doc) for i, doc in enumerate(documents)) ) # exact-content dedup (independent of `deduplicate`): pre-fetch @@ -740,14 +754,17 @@ async def create_documents( if doc.embedding: docs_with_embeddings.append((new_doc, doc.embedding)) + except IntegrityError as e: + await db.rollback() + raise ValidationException( + "Failed to create documents due to integrity constraint violation" + ) from e except SQLAlchemyError: - # The session/transaction is dead; continuing the loop would only - # cascade PendingRollbackErrors and lose the whole batch silently. + # Dead transaction: continuing would cascade PendingRollbackErrors. await db.rollback() raise except Exception as e: - # Genuinely per-document failures (bad content, metadata, token - # overflow) skip the document without poisoning the batch. + # Per-document failures (bad content, metadata, token overflow). logger.error( f"Error adding new document to {workspace_name}/{doc.session_name}/{observer}/{observed}: {e}" ) @@ -1304,7 +1321,7 @@ async def _semantic_dup_decision( document_ids=candidate_document_ids, filters=filters, ) - else: + elif _uses_pgvector(): similar_docs = await query_documents( db=db, workspace_name=workspace_name, @@ -1316,6 +1333,8 @@ async def _semantic_dup_decision( top_k=_SEMANTIC_DUP_TOP_K, embedding=doc.embedding or None, ) + else: + return SemanticRejectionResult.NOT_DUPLICATE, None if not similar_docs: return SemanticRejectionResult.NOT_DUPLICATE, None diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index cc19461a..ade79c80 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -54,10 +54,12 @@ logger = getLogger(__name__) load_dotenv(override=True) -# Total processing attempts per work unit for transient errors, counted -# per instance (N deriver instances give N x this many attempts). +# Total processing attempts per work unit for transient errors. Count is +# stored on the oldest unprocessed queue item so every deriver instance +# shares one budget. MAX_RETRYABLE_ATTEMPTS = 3 RETRY_BACKOFF_SECONDS = 1.0 +_RETRY_ATTEMPTS_PAYLOAD_KEY = "_retry_attempts" class WorkerOwnership(NamedTuple): @@ -135,10 +137,6 @@ class QueueManager: self.worker_ownership: dict[str, WorkerOwnership] = {} self.queue_empty_flag: asyncio.Event = asyncio.Event() - # Transient-failure attempts per work-unit key. Entries are removed - # on success or terminal failure, never on the retry path itself. - self._retry_attempts: dict[str, int] = {} - # Current adaptive polling interval; grows while idle/erroring and # resets to the base interval as soon as work is claimed. self._current_poll_interval: float = ( @@ -596,6 +594,8 @@ class QueueManager: Transient errors (is_retryable_error) get up to MAX_RETRYABLE_ATTEMPTS attempts per work unit: items stay unprocessed with no error recorded. + The attempt count lives on the oldest unprocessed queue item so a + different deriver instance continues the same budget after reclaim. Reprocessing is safe because a retried batch re-derives the same observations and exact dedup collapses them into reinforcement. @@ -610,9 +610,9 @@ class QueueManager: context: Context string describing what was being processed (e.g., "processing representation batch") """ if is_retryable_error(error): - attempts = self._retry_attempts.get(work_unit_key, 0) + 1 + attempts = await self._get_work_unit_retry_attempts(work_unit_key) + 1 if attempts < MAX_RETRYABLE_ATTEMPTS: - self._retry_attempts[work_unit_key] = attempts + await self._set_work_unit_retry_attempts(work_unit_key, attempts) logger.warning( "Transient error %s for work unit %s (attempt %d/%d); leaving items unprocessed for retry", context, @@ -623,7 +623,7 @@ class QueueManager: ) return True - self._retry_attempts.pop(work_unit_key, None) + await self._clear_work_unit_retry_attempts(work_unit_key) error_msg = f"{error.__class__.__name__}: {str(error)}" try: if items: @@ -706,7 +706,9 @@ class QueueManager: await self.mark_queue_items_as_processed( items_to_process, work_unit_key ) - self._retry_attempts.pop(work_unit_key, None) + await self._clear_work_unit_retry_attempts( + work_unit_key + ) queue_item_count += len(items_to_process) except Exception as e: if await self._handle_processing_error( @@ -737,7 +739,9 @@ class QueueManager: await self.mark_queue_items_as_processed( [queue_item], work_unit_key ) - self._retry_attempts.pop(work_unit_key, None) + await self._clear_work_unit_retry_attempts( + work_unit_key + ) queue_item_count += 1 except Exception as e: if await self._handle_processing_error( @@ -1102,6 +1106,83 @@ class QueueManager: batch_max_tokens=batch_max_tokens, ) + async def _oldest_unprocessed_item( + self, + db: AsyncSession, + work_unit_key: str, + *, + for_update: bool = False, + ) -> models.QueueItem | None: + stmt = ( + select(models.QueueItem) + .where( + models.QueueItem.work_unit_key == work_unit_key, + models.QueueItem.processed.is_(False), + ) + .order_by(models.QueueItem.id) + .limit(1) + ) + if for_update: + stmt = stmt.with_for_update() + result = await db.execute(stmt) + return result.scalar_one_or_none() + + async def _get_work_unit_retry_attempts(self, work_unit_key: str) -> int: + """Read the shared transient-failure attempt count for a work unit.""" + async with tracked_db("get_work_unit_retry_attempts") as db: + item = await self._oldest_unprocessed_item(db, work_unit_key) + if item is None: + return 0 + raw = (item.payload or {}).get(_RETRY_ATTEMPTS_PAYLOAD_KEY, 0) + try: + return max(0, int(raw)) + except (TypeError, ValueError): + return 0 + + async def _set_work_unit_retry_attempts( + self, work_unit_key: str, attempts: int + ) -> None: + """Persist the shared attempt count on the oldest unprocessed item.""" + async with tracked_db("set_work_unit_retry_attempts") as db: + item = await self._oldest_unprocessed_item( + db, work_unit_key, for_update=True + ) + if item is None: + await db.commit() + return + new_payload = dict(item.payload or {}) + new_payload[_RETRY_ATTEMPTS_PAYLOAD_KEY] = attempts + await db.execute( + update(models.QueueItem) + .where(models.QueueItem.id == item.id) + .values(payload=new_payload) + ) + await db.commit() + + async def _clear_work_unit_retry_attempts(self, work_unit_key: str) -> None: + """Drop the shared attempt count after success or terminal failure.""" + async with tracked_db("clear_work_unit_retry_attempts") as db: + result = await db.execute( + select(models.QueueItem) + .where( + models.QueueItem.work_unit_key == work_unit_key, + models.QueueItem.processed.is_(False), + ) + .with_for_update() + ) + for item in result.scalars(): + payload = item.payload or {} + if _RETRY_ATTEMPTS_PAYLOAD_KEY not in payload: + continue + new_payload = dict(payload) + new_payload.pop(_RETRY_ATTEMPTS_PAYLOAD_KEY, None) + await db.execute( + update(models.QueueItem) + .where(models.QueueItem.id == item.id) + .values(payload=new_payload) + ) + await db.commit() + async def mark_queue_items_as_processed( self, items: list[QueueItem], work_unit_key: str ) -> None: diff --git a/tests/crud/test_document.py b/tests/crud/test_document.py index 2e4e4e94..b9699fe4 100644 --- a/tests/crud/test_document.py +++ b/tests/crud/test_document.py @@ -1745,3 +1745,45 @@ class TestExternalCandidateHoist: assert len(result.created_documents) == 2 assert events[:2] == ["resolve", "resolve"] assert "execute" in events + + @pytest.mark.asyncio + async def test_resolve_failure_skips_semantic_without_query_documents( + self, + db_session: AsyncSession, + sample_data: tuple[models.Workspace, models.Peer], + monkeypatch: pytest.MonkeyPatch, + ): + from src.config import settings + + test_workspace, test_peer = sample_data + observed_peer, test_session = await self._setup( + db_session, test_workspace, test_peer + ) + monkeypatch.setattr(settings.VECTOR_STORE, "TYPE", "turbopuffer") + monkeypatch.setattr(settings.VECTOR_STORE, "MIGRATED", True) + + with ( + patch( + "src.crud.document.query_external_vector_document_ids", + side_effect=RuntimeError("store down"), + ), + patch( + "src.crud.document.get_external_vector_store", + return_value=None, + ), + patch( + "src.crud.document.query_documents", + new_callable=AsyncMock, + ) as mock_query, + ): + result = await crud.create_documents( + db_session, + [self._doc("fact one", test_session.name)], + workspace_name=test_workspace.name, + observer=test_peer.name, + observed=observed_peer.name, + deduplicate=True, + ) + + assert len(result.created_documents) == 1 + mock_query.assert_not_awaited() diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index f37067ff..933299ae 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -1986,6 +1986,16 @@ class TestQueueRetry: .all() ) + async def _retry_attempts_on_items( + self, db_session: AsyncSession, work_unit_key: str + ) -> int | None: + items = await self._fetch_items(db_session, work_unit_key) + unprocessed = [item for item in items if not item.processed] + if not unprocessed: + return None + raw = (unprocessed[0].payload or {}).get("_retry_attempts") + return None if raw is None else int(raw) + async def test_retryable_error_leaves_items_unprocessed( self, db_session: AsyncSession, @@ -2024,7 +2034,7 @@ class TestQueueRetry: assert all(not item.processed for item in items) assert all(item.error is None for item in items) assert await self._aqs_rows(db_session, work_unit_key) == 0 - assert qm._retry_attempts[work_unit_key] == 1 # pyright: ignore[reportPrivateUsage] + assert await self._retry_attempts_on_items(db_session, work_unit_key) == 1 assert qm.semaphore._value == initial_semaphore_value async def test_retry_exhaustion_is_terminal( @@ -2042,7 +2052,9 @@ class TestQueueRetry: qm, work_unit_key, worker_id, _ = await self._seed_work_unit( db_session, sample_session_with_peers, create_queue_payload ) - qm._retry_attempts[work_unit_key] = MAX_RETRYABLE_ATTEMPTS - 1 # pyright: ignore[reportPrivateUsage] + await qm._set_work_unit_retry_attempts( # pyright: ignore[reportPrivateUsage] + work_unit_key, MAX_RETRYABLE_ATTEMPTS - 1 + ) with patch( "src.deriver.queue_manager.process_representation_batch", @@ -2055,7 +2067,7 @@ class TestQueueRetry: assert items[0].processed assert items[0].error is not None assert "OperationalError" in items[0].error - assert work_unit_key not in qm._retry_attempts # pyright: ignore[reportPrivateUsage] + assert await self._retry_attempts_on_items(db_session, work_unit_key) is None async def test_non_retryable_error_burns_immediately( self, @@ -2080,7 +2092,7 @@ class TestQueueRetry: assert items[0].processed assert items[0].error is not None assert "ValueError" in items[0].error - assert work_unit_key not in qm._retry_attempts # pyright: ignore[reportPrivateUsage] + assert await self._retry_attempts_on_items(db_session, work_unit_key) is None async def test_counter_cleared_after_success( self, @@ -2092,7 +2104,7 @@ class TestQueueRetry: qm, work_unit_key, worker_id, _ = await self._seed_work_unit( db_session, sample_session_with_peers, create_queue_payload ) - qm._retry_attempts[work_unit_key] = 1 # pyright: ignore[reportPrivateUsage] + await qm._set_work_unit_retry_attempts(work_unit_key, 1) # pyright: ignore[reportPrivateUsage] async def noop_batch(*_args: Any, **_kwargs: Any) -> None: return None @@ -2106,4 +2118,52 @@ class TestQueueRetry: items = await self._fetch_items(db_session, work_unit_key) assert all(item.processed for item in items) assert all(item.error is None for item in items) - assert work_unit_key not in qm._retry_attempts # pyright: ignore[reportPrivateUsage] + assert await self._retry_attempts_on_items(db_session, work_unit_key) is None + + async def test_retry_budget_survives_reclaim_by_another_manager( + self, + db_session: AsyncSession, + sample_session_with_peers: tuple[models.Session, list[models.Peer]], + create_queue_payload: Callable[..., Any], + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """A second QueueManager continues the durable attempt budget.""" + from src.deriver.queue_manager import MAX_RETRYABLE_ATTEMPTS + + monkeypatch.setattr("src.deriver.queue_manager.RETRY_BACKOFF_SECONDS", 0.0) + qm1, work_unit_key, worker_id, _ = await self._seed_work_unit( + db_session, sample_session_with_peers, create_queue_payload + ) + + with patch( + "src.deriver.queue_manager.process_representation_batch", + side_effect=self._retryable_error(), + ): + await qm1.process_work_unit(work_unit_key, worker_id) + + assert await self._retry_attempts_on_items(db_session, work_unit_key) == 1 + assert await self._aqs_rows(db_session, work_unit_key) == 0 + + # Seed the remaining budget so the next reclaim is the terminal attempt. + qm2 = QueueManager() + await qm2._set_work_unit_retry_attempts( # pyright: ignore[reportPrivateUsage] + work_unit_key, MAX_RETRYABLE_ATTEMPTS - 1 + ) + claimed = await qm2.claim_work_units(db_session, [work_unit_key]) + worker_id_2 = "test_worker_2" + qm2.worker_ownership[worker_id_2] = WorkerOwnership( + work_unit_key=work_unit_key, aqs_id=claimed[work_unit_key] + ) + await db_session.commit() + + with patch( + "src.deriver.queue_manager.process_representation_batch", + side_effect=self._retryable_error(), + ): + await qm2.process_work_unit(work_unit_key, worker_id_2) + + items = await self._fetch_items(db_session, work_unit_key) + assert len(items) == 1 + assert items[0].processed + assert items[0].error is not None + assert "OperationalError" in items[0].error