diff --git a/apps/api/alembic/versions/4a5b6c7d8e9f_add_retrieval_serving_generations.py b/apps/api/alembic/versions/4a5b6c7d8e9f_add_retrieval_serving_generations.py index 426ff6e0..5972f228 100644 --- a/apps/api/alembic/versions/4a5b6c7d8e9f_add_retrieval_serving_generations.py +++ b/apps/api/alembic/versions/4a5b6c7d8e9f_add_retrieval_serving_generations.py @@ -13,6 +13,15 @@ branch_labels: Sequence[str] | None = None depends_on: Sequence[str] | None = None +__all__: list[str] = [ + "revision", + "down_revision", + "branch_labels", + "depends_on", + "upgrade", + "downgrade", +] + def upgrade() -> None: if not sa.inspect(op.get_bind()).has_table("retrieval_namespace_generations"): diff --git a/apps/api/alembic/versions/5b6c7d8e9f0a_add_term_trigram_indexes.py b/apps/api/alembic/versions/5b6c7d8e9f0a_add_term_trigram_indexes.py index 71720ff0..50126978 100644 --- a/apps/api/alembic/versions/5b6c7d8e9f0a_add_term_trigram_indexes.py +++ b/apps/api/alembic/versions/5b6c7d8e9f0a_add_term_trigram_indexes.py @@ -12,6 +12,15 @@ branch_labels: Sequence[str] | None = None depends_on: Sequence[str] | None = None +__all__: list[str] = [ + "revision", + "down_revision", + "branch_labels", + "depends_on", + "upgrade", + "downgrade", +] + _MAP_UNIT_INDEX = "idx_document_map_units_term_trgm" _CHUNK_INDEX = "idx_document_chunks_term_trgm" diff --git a/apps/api/alembic/versions/6c7d8e9f0a1b_add_retrieval_namespace_map_snapshots.py b/apps/api/alembic/versions/6c7d8e9f0a1b_add_retrieval_namespace_map_snapshots.py index 1f7b7c01..b561548f 100644 --- a/apps/api/alembic/versions/6c7d8e9f0a1b_add_retrieval_namespace_map_snapshots.py +++ b/apps/api/alembic/versions/6c7d8e9f0a1b_add_retrieval_namespace_map_snapshots.py @@ -13,6 +13,15 @@ branch_labels: Sequence[str] | None = None depends_on: Sequence[str] | None = None +__all__: list[str] = [ + "revision", + "down_revision", + "branch_labels", + "depends_on", + "upgrade", + "downgrade", +] + def upgrade() -> None: if not sa.inspect(op.get_bind()).has_table("retrieval_namespace_map_snapshots"): diff --git a/apps/api/tests/contract/test_retrieval_snapshot_large_corpus_contract.py b/apps/api/tests/contract/test_retrieval_snapshot_large_corpus_contract.py index 0b45a703..f200ff84 100644 --- a/apps/api/tests/contract/test_retrieval_snapshot_large_corpus_contract.py +++ b/apps/api/tests/contract/test_retrieval_snapshot_large_corpus_contract.py @@ -16,7 +16,7 @@ _REVISION_GROUP_SIZE, load_nav_snapshot, ) -import shared.services.retrieval.nav_snapshot as nav_snapshot_module +from shared.services.retrieval import nav_snapshot as nav_snapshot_module from sqlalchemy import Executable, Result, select from sqlalchemy.engine import Row from sqlalchemy.sql.selectable import Select @@ -64,6 +64,9 @@ async def execute(self, statement: Executable) -> Result[tuple[object, ...]]: result = await self._session.execute(statement) return cast(Result[tuple[object, ...]], result) + async def rollback(self) -> None: + await self._session.rollback() + async def _seed_large_retrieval_corpus(namespace: str) -> None: await ContractDatabase.execute( diff --git a/packages/shared-python/shared/services/retrieval/execution/routes.py b/packages/shared-python/shared/services/retrieval/execution/routes.py index 673e2b6e..29973c6b 100644 --- a/packages/shared-python/shared/services/retrieval/execution/routes.py +++ b/packages/shared-python/shared/services/retrieval/execution/routes.py @@ -56,7 +56,6 @@ async def run_retrieval_route( async def _try_run_small_corpus_route( context: RetrievalRouteContext, ) -> RetrievalRouteOutcome | None: - total_chunk_count: int | None = None total_chunk_count = await count_scoped_chunks( context.db, user_id=context.user_id, diff --git a/packages/shared-python/shared/services/retrieval/nav_snapshot.py b/packages/shared-python/shared/services/retrieval/nav_snapshot.py index 95d24078..ea502c54 100644 --- a/packages/shared-python/shared/services/retrieval/nav_snapshot.py +++ b/packages/shared-python/shared/services/retrieval/nav_snapshot.py @@ -26,6 +26,7 @@ tuple_, ) from sqlalchemy.engine import Result +from sqlalchemy.exc import SQLAlchemyError from shared.models.database.document import ( Document, @@ -393,8 +394,14 @@ async def _resolve_namespace_snapshot_entries( ) try: row = (await db.execute(statement)).first() - except Exception: + except SQLAlchemyError as exc: await db.rollback() + _logger.warning( + "retrieval snapshot namespace lookup failed user_id=%s namespace=%s error=%s", + user_id, + namespace, + exc, + ) return None if row is None: return None @@ -470,8 +477,13 @@ async def _resolve_manifest_entries( ) try: rows = (await db.execute(statement)).all() - except Exception: + except SQLAlchemyError as exc: await db.rollback() + _logger.warning( + "retrieval manifest lookup failed revisions=%s error=%s", + revision_group, + exc, + ) return None if len(rows) != len(revision_group): return None