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
@@ -0,0 +1,52 @@
"""Add demo materialization claim state."""

from __future__ import annotations

from collections.abc import Sequence

from alembic import op
import sqlalchemy as sa


revision: str = "0b1c2d3e4f5a"
down_revision: str | Sequence[str] | None = "e4f5a6b7c8d9"
branch_labels: Sequence[str] | None = None
depends_on: Sequence[str] | None = None


def upgrade() -> None:
op.add_column(
"demo_materializations",
sa.Column("status", sa.String(length=32), nullable=True),
)
op.add_column(
"demo_materializations",
sa.Column("claimed_at", sa.DateTime(), nullable=True),
)
op.execute(
"UPDATE demo_materializations SET status = 'ready' WHERE status IS NULL"
)
op.alter_column(
"demo_materializations",
"status",
existing_type=sa.String(length=32),
nullable=False,
server_default="ready",
)
op.alter_column(
"demo_materializations",
"document_id",
existing_type=sa.String(length=36),
nullable=True,
)


def downgrade() -> None:
op.alter_column(
"demo_materializations",
"document_id",
existing_type=sa.String(length=36),
nullable=False,
)
op.drop_column("demo_materializations", "claimed_at")
op.drop_column("demo_materializations", "status")
222 changes: 146 additions & 76 deletions apps/api/app/services/demo/source_materializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,23 +4,25 @@

import shutil
import tempfile
from collections.abc import Iterable
from dataclasses import dataclass
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from hashlib import blake2b
from pathlib import Path
from uuid import uuid4

from app.services.demo.source_catalog import DemoSourceCatalog, DemoSourceDefinition
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy import delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession

from shared.core.exceptions.domain_exceptions import ValidationException
from shared.core.exceptions.domain_exceptions import ConflictException, ValidationException
from shared.models.database.demo_materialization import DemoMaterialization
from shared.models.database.document import Document
from shared.models.database.job import Job
from shared.models.database.job_result import JobResult
from shared.services.retrieval.cache_service import invalidate_retrieval_cache_namespaces
from shared.services.retrieval.publication_service import RetrievalPublicationService
from shared.services.retrieval.publication_models import DocumentPublicationScope
from shared.services.storage.result_storage import get_result_storage


Expand Down Expand Up @@ -73,17 +75,39 @@ async def materialize_sources(
self._catalog.require_source(demo_source_id)
for demo_source_id in selected_demo_source_ids
]
results: list[MaterializedDemoSource] = []
for source in selected_sources:
result = await self._materialize_source(
try:
claims = await self._claim_sources(
db,
user_id=user_id,
namespace=namespace,
source=source,
sources=selected_sources,
)
results.append(result)

except IntegrityError as error:
await db.rollback()
raise ConflictException(
user_message="This demo source is currently being materialized.",
reason="ABORTED",
resource="Demo materialization",
) from error
await db.commit()
results: list[MaterializedDemoSource] = []
for source in selected_sources:
try:
result = await self._materialize_source(
db,
user_id=user_id,
namespace=namespace,
source=source,
claim=claims[source.demo_source_id],
)
results.append(result)
except Exception:
await db.rollback()
await self._release_claims(
db,
claims=claims.values(),
)
raise
await invalidate_retrieval_cache_namespaces(
user_id=user_id,
namespaces=[namespace],
Expand All @@ -97,29 +121,8 @@ async def _materialize_source(
user_id: str,
namespace: str,
source: DemoSourceDefinition,
claim: DemoMaterialization,
) -> MaterializedDemoSource:
await _lock_materialization_scope(
db,
user_id=user_id,
namespace=namespace,
demo_source_id=source.demo_source_id,
)
existing = await self._get_existing_materialization(
db,
user_id=user_id,
namespace=namespace,
demo_source_id=source.demo_source_id,
)
if existing is not None and await self._is_active_document(
db,
document_id=existing.document_id,
):
return _materialized_source_payload(
source=source,
document_id=existing.document_id,
status="existing",
)

document_id = f"doc_{uuid4().hex[:12]}"
job_id = f"job_demo_{uuid4().hex[:12]}"
job_result_id = str(uuid4())
Expand Down Expand Up @@ -169,12 +172,13 @@ async def _materialize_source(
)
await db.flush()
chunks = self._catalog.publication_chunks(source)
await db.run_sync(
published_state = await db.run_sync(
lambda sync_db: self._publication_service.publish_document_state(
sync_db,
job_id=job_id,
job_result_id=job_result_id,
chunks=[dict(chunk) for chunk in chunks],
update_namespace_snapshot=False,
)
)
await db.run_sync(
Expand All @@ -186,58 +190,114 @@ async def _materialize_source(
)
await db.flush()

if existing is None:
db.add(
DemoMaterialization(
if published_state is None or published_state.document_id != document_id:
raise RuntimeError("Demo publication did not create its requested document")
if published_state.manifest_payload is None:
raise RuntimeError("Demo publication did not create a serving manifest")
manifest_payload = published_state.manifest_payload
await db.run_sync(
lambda sync_db: self._publication_service.update_namespace_snapshot(
sync_db,
scope=DocumentPublicationScope(
user_id=user_id,
namespace=namespace,
demo_source_id=source.demo_source_id,
document_id=document_id,
created_at=timestamp,
updated_at=timestamp,
)
job_result_id=job_result_id,
source_file_name=source.title,
),
manifest_payload=manifest_payload,
)
else:
existing.document_id = document_id
existing.updated_at = timestamp
)
await db.commit()

claim.document_id = document_id
claim.status = "ready"
claim.claimed_at = None
claim.updated_at = timestamp
await db.flush()
await db.commit()
return _materialized_source_payload(
source=source,
document_id=document_id,
status="created",
)

async def _get_existing_materialization(
async def _claim_sources(
self,
db: AsyncSession,
*,
user_id: str,
namespace: str,
demo_source_id: str,
) -> DemoMaterialization | None:
result = await db.execute(
select(DemoMaterialization)
.where(DemoMaterialization.user_id == user_id)
.where(DemoMaterialization.namespace == namespace)
.where(DemoMaterialization.demo_source_id == demo_source_id)
.with_for_update()
.limit(1)
)
return result.scalar_one_or_none()
sources: list[DemoSourceDefinition],
) -> dict[str, DemoMaterialization]:
now = _utc_now()
claims: dict[str, DemoMaterialization] = {}
for source in sorted(sources, key=lambda item: item.demo_source_id):
lock_acquired = await db.scalar(
select(
func.pg_try_advisory_xact_lock(
_materialization_lock_id(
user_id=user_id,
namespace=namespace,
demo_source_id=source.demo_source_id,
)
)
)
)
if not lock_acquired:
raise ConflictException(
user_message="This demo source is currently being materialized.",
reason="ABORTED",
resource="Demo materialization",
resource_id=source.demo_source_id,
)
result = await db.execute(
select(DemoMaterialization)
.where(DemoMaterialization.user_id == user_id)
.where(DemoMaterialization.namespace == namespace)
.where(DemoMaterialization.demo_source_id == source.demo_source_id)
)
existing = result.scalar_one_or_none()
if existing is not None:
if not _is_stale_claim(existing, now=now):
raise _materialization_conflict(existing, source)
claim = existing
claim.status = "materializing"
claim.document_id = None
claim.claimed_at = now
claim.updated_at = now
else:
claim = DemoMaterialization(
user_id=user_id,
namespace=namespace,
demo_source_id=source.demo_source_id,
status="materializing",
document_id=None,
claimed_at=now,
created_at=now,
updated_at=now,
)
db.add(claim)
claims[source.demo_source_id] = claim
await db.flush()
return claims

async def _is_active_document(
async def _release_claims(
self,
db: AsyncSession,
*,
document_id: str,
) -> bool:
result = await db.execute(
select(Document.document_id)
.where(Document.document_id == document_id)
.where(Document.status == "active")
.limit(1)
claims: Iterable[DemoMaterialization],
) -> None:
claim_ids = [claim.id for claim in claims]
if not claim_ids:
return
await db.execute(
delete(DemoMaterialization).where(
DemoMaterialization.status == "materializing",
DemoMaterialization.id.in_(claim_ids),
)
)
return result.scalar_one_or_none() is not None
await db.commit()


def _deduplicate_source_ids(demo_source_ids: list[str]) -> list[str]:
Expand All @@ -252,19 +312,12 @@ def _deduplicate_source_ids(demo_source_ids: list[str]) -> list[str]:
return selected


async def _lock_materialization_scope(
db: AsyncSession,
*,
user_id: str,
namespace: str,
demo_source_id: str,
) -> None:
lock_id = _materialization_lock_id(
user_id=user_id,
namespace=namespace,
demo_source_id=demo_source_id,
def _is_stale_claim(materialization: DemoMaterialization, *, now: datetime) -> bool:
return (
materialization.status == "materializing"
and materialization.claimed_at is not None
and materialization.claimed_at < now - timedelta(minutes=10)
)
await db.execute(select(func.pg_advisory_xact_lock(lock_id)))


def _materialization_lock_id(
Expand All @@ -278,6 +331,23 @@ def _materialization_lock_id(
return int.from_bytes(digest, byteorder="big", signed=True)


def _materialization_conflict(
materialization: DemoMaterialization,
source: DemoSourceDefinition,
) -> ConflictException:
reason = "ALREADY_EXISTS" if materialization.status == "ready" else "ABORTED"
return ConflictException(
user_message=(
"This demo source has already been materialized."
if reason == "ALREADY_EXISTS"
else "This demo source is currently being materialized."
),
reason=reason,
resource="Demo materialization",
resource_id=source.demo_source_id,
)


def _materialized_source_payload(
*,
source: DemoSourceDefinition,
Expand Down
Loading
Loading