Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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"):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
16 changes: 14 additions & 2 deletions packages/shared-python/shared/services/retrieval/nav_snapshot.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
tuple_,
)
from sqlalchemy.engine import Result
from sqlalchemy.exc import SQLAlchemyError

from shared.models.database.document import (
Document,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading