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
122 changes: 88 additions & 34 deletions apps/api/app/services/demo/source_materializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,16 @@

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

import logfire

from app.services.demo.source_catalog import DemoSourceCatalog, DemoSourceDefinition
from sqlalchemy.exc import IntegrityError
from sqlalchemy import delete, func, select
Expand All @@ -23,6 +26,8 @@
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.redis import RedisPublicationSemaphore, RedisServiceFactory
from shared.core.config import settings
from shared.services.storage.result_storage import get_result_storage


Expand Down Expand Up @@ -50,6 +55,7 @@ def __init__(
) -> None:
self._catalog = catalog
self._publication_service = publication_service or RetrievalPublicationService()
self._redis_service = RedisServiceFactory.get_service()

async def materialize_sources(
self,
Expand Down Expand Up @@ -93,12 +99,19 @@ async def materialize_sources(
results: list[MaterializedDemoSource] = []
for source in selected_sources:
try:
semaphore = RedisPublicationSemaphore(
self._redis_service,
concurrency=settings.MATERIALIZATION_DB_PUBLICATION_CONCURRENCY,
lease_seconds=settings.MATERIALIZATION_DB_PUBLICATION_LEASE_SECONDS,
acquire_timeout_seconds=settings.MATERIALIZATION_DB_PUBLICATION_ACQUIRE_TIMEOUT_SECONDS,
)
result = await self._materialize_source(
db,
user_id=user_id,
namespace=namespace,
source=source,
claim=claims[source.demo_source_id],
semaphore=semaphore,
)
results.append(result)
except Exception:
Expand All @@ -122,15 +135,22 @@ async def _materialize_source(
namespace: str,
source: DemoSourceDefinition,
claim: DemoMaterialization,
semaphore: RedisPublicationSemaphore,
) -> MaterializedDemoSource:
document_id = f"doc_{uuid4().hex[:12]}"
job_id = f"job_demo_{uuid4().hex[:12]}"
job_result_id = str(uuid4())
timestamp = _utc_now()
stage_started_at = time.perf_counter()
result_bundle = _upload_demo_result_bundle(
job_id=job_id,
source_directory=self._catalog.source_directory(source),
)
logfire.info(
"Demo materialization source bundle upload completed",
demo_source_id=source.demo_source_id,
duration_seconds=time.perf_counter() - stage_started_at,
)

db.add(
Job(
Expand Down Expand Up @@ -170,45 +190,79 @@ async def _materialize_source(
updated_at=timestamp,
)
)
await db.flush()
chunks = self._catalog.publication_chunks(source)
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,
)
stage_started_at = time.perf_counter()
wait_seconds = await semaphore.acquire()
logfire.info(
"Demo materialization publication semaphore acquired",
demo_source_id=source.demo_source_id,
wait_seconds=wait_seconds,
)
await db.run_sync(
lambda sync_db: self._publication_service.publish_document_graph(
sync_db,
job_id=job_id,
job_result_id=job_result_id,
try:
base_rows_started_at = time.perf_counter()
await db.flush()
logfire.info(
"Demo materialization base rows completed",
demo_source_id=source.demo_source_id,
duration_seconds=time.perf_counter() - base_rows_started_at,
)
)
await db.flush()

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,
document_id=document_id,
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,
source_file_name=source.title,
),
manifest_payload=manifest_payload,
chunks=[dict(chunk) for chunk in chunks],
update_namespace_snapshot=False,
)
)
)
await db.commit()
await db.run_sync(
lambda sync_db: self._publication_service.publish_document_graph(
sync_db,
job_id=job_id,
job_result_id=job_result_id,
)
)
await db.flush()
logfire.info(
"Demo materialization sections chunks and map index completed",
demo_source_id=source.demo_source_id,
duration_seconds=time.perf_counter() - stage_started_at,
chunk_count=len(chunks),
)

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
stage_started_at = time.perf_counter()
await db.run_sync(
lambda sync_db: self._publication_service.update_namespace_snapshot(
sync_db,
scope=DocumentPublicationScope(
user_id=user_id,
namespace=namespace,
document_id=document_id,
job_result_id=job_result_id,
source_file_name=source.title,
),
manifest_payload=manifest_payload,
)
)
logfire.info(
"Demo materialization namespace snapshot completed",
demo_source_id=source.demo_source_id,
duration_seconds=time.perf_counter() - stage_started_at,
)
commit_started_at = time.perf_counter()
await db.commit()
logfire.info(
"Demo materialization publication commit completed",
demo_source_id=source.demo_source_id,
duration_seconds=time.perf_counter() - commit_started_at,
)
finally:
await semaphore.release()

claim.document_id = document_id
claim.status = "ready"
Expand Down
16 changes: 16 additions & 0 deletions packages/shared-python/shared/core/config/database.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,22 @@ class DatabaseConfig(BaseModel):
default=50, description="Celery gevent worker concurrency"
)

MATERIALIZATION_DB_PUBLICATION_CONCURRENCY: int = Field(
default=2,
ge=1,
description="Maximum concurrent materialization database publications",
)
MATERIALIZATION_DB_PUBLICATION_LEASE_SECONDS: int = Field(
default=900,
ge=60,
description="Redis lease duration for a materialization publication permit",
)
MATERIALIZATION_DB_PUBLICATION_ACQUIRE_TIMEOUT_SECONDS: float = Field(
default=30.0,
ge=0.1,
description="Maximum time to wait for a materialization publication permit",
)

def get_ssl_connect_args(self) -> dict:
"""Return SSL connect args for psycopg2."""
ssl_args = {"sslmode": self.DB_SSL_MODE}
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

from contextlib import nullcontext
from dataclasses import dataclass
from typing import Any

Expand All @@ -11,6 +12,8 @@
from shared.models.schemas.job_metadata import JobMetadataHelper
from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace
from shared.services.redis.redis_sync_service import SyncRedisServiceFactory
from shared.services.redis.publication_semaphore import SyncRedisPublicationSemaphore
from shared.core.config import settings
from shared.services.retrieval.publication_service import RetrievalPublicationService
from shared.services.retrieval.publication_models import (
ExistingDocumentScope,
Expand Down Expand Up @@ -53,25 +56,38 @@ def publish_result(
section_summaries: dict[str, str] | None,
document_top_summary: str | None = None,
) -> JobPublicationOutcome:
previous_document_scope = self._retrieval_publication.get_existing_document_scope(
db,
job_id=job_id,
semaphore = SyncRedisPublicationSemaphore(
SyncRedisServiceFactory.get_service(),
concurrency=settings.MATERIALIZATION_DB_PUBLICATION_CONCURRENCY,
lease_seconds=settings.MATERIALIZATION_DB_PUBLICATION_LEASE_SECONDS,
acquire_timeout_seconds=settings.MATERIALIZATION_DB_PUBLICATION_ACQUIRE_TIMEOUT_SECONDS,
)
published_document_state = self._retrieval_publication.publish_document_state(
db,
job_id=job_id,
job_result_id=job_result_id,
chunks=chunks,
section_summaries=section_summaries,
job_type = db.execute(
select(Job.job_type).where(Job.job_id == job_id)
).scalar_one_or_none()
publication_context = (
semaphore if job_type == "demo_materialization" else nullcontext()
)
if _should_publish_document_graph(published_document_state):
assert published_document_state is not None
self._retrieval_publication.publish_document_graph(
with publication_context:
previous_document_scope = self._retrieval_publication.get_existing_document_scope(
db,
job_id=job_id,
)
published_document_state = self._retrieval_publication.publish_document_state(
db,
job_id=job_id,
job_result_id=job_result_id,
top_summary=document_top_summary,
chunks=chunks,
section_summaries=section_summaries,
)
if _should_publish_document_graph(published_document_state):
assert published_document_state is not None
self._retrieval_publication.publish_document_graph(
db,
job_id=job_id,
job_result_id=job_result_id,
top_summary=document_top_summary,
)

cache_invalidation = self._build_cache_invalidation(
db,
Expand Down
8 changes: 8 additions & 0 deletions packages/shared-python/shared/services/redis/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,10 @@
from .key_builder import RedisKeyBuilder, RedisKeyType, redis_key_builder
from .redis_alerts import AlertRule, RedisAlertManager, RedisAlertNotifier
from .redis_monitor import RedisMonitor
from .publication_semaphore import (
RedisPublicationSemaphore,
SyncRedisPublicationSemaphore,
)
from .redis_service import RedisService
from .redis_service_factory import RedisServiceFactory
from .retry_policy import RedisHealthChecker, RedisRetry
Expand All @@ -20,6 +24,8 @@
__all__ = [
"RedisService",
"RedisServiceFactory",
"RedisPublicationSemaphore",
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
"SyncRedisPublicationSemaphore",
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
"RedisMonitor",
"RedisAlertManager",
"RedisAlertNotifier",
Expand All @@ -38,6 +44,8 @@
_EXPORT_MODULES: dict[str, str] = {
"RedisService": "shared.services.redis.redis_service",
"RedisServiceFactory": "shared.services.redis.redis_service_factory",
"RedisPublicationSemaphore": "shared.services.redis.publication_semaphore",
"SyncRedisPublicationSemaphore": "shared.services.redis.publication_semaphore",
"RedisMonitor": "shared.services.redis.redis_monitor",
"RedisAlertManager": "shared.services.redis.redis_alerts",
"RedisAlertNotifier": "shared.services.redis.redis_alerts",
Expand Down
Loading
Loading