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
13 changes: 9 additions & 4 deletions src/core/ingestion/application/ingestion_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -538,10 +538,13 @@ async def process_document(self, document_id: str, force: bool = False) -> None:
)
return
generation = None
# A failed flush may expire ORM attributes before explicit rollback.
generation_id: str | None = None
try:
if pending_generation_id:
generation = await self.document_repository.get_generation(pending_generation_id)
if generation is not None:
generation_id = generation.id
# A retry (re-upload with the same content-hash, or a
# stale-lock sweep) reaches here with pending_generation_id
# still pointing at the PREVIOUS attempt's generation row.
Expand Down Expand Up @@ -577,6 +580,7 @@ async def process_document(self, document_id: str, force: bool = False) -> None:
keywords=list(getattr(document, "keywords", None) or []),
hashtags=list(getattr(document, "hashtags", None) or []),
)
generation_id = generation.id
await self.document_repository.save_generation(generation)
document.pending_generation_id = generation.id
await self.document_repository.save(document)
Expand Down Expand Up @@ -1296,11 +1300,12 @@ async def _on_graph_progress(completed: int, total: int):
logger.error(f"Failed to map error for {document_id}: {map_err}")
error_message = f"{type(e).__name__}: {str(e)}"

await self.document_repository.mark_generation_failed(
generation.id, error_message
)
if generation_id:
await self.document_repository.mark_generation_failed(
generation_id, error_message
)
if preserve_published:
if document.pending_generation_id == generation.id:
if document.pending_generation_id == generation_id:
document.pending_generation_id = None
await self.document_repository.save(document)
else:
Expand Down
128 changes: 128 additions & 0 deletions tests/unit/test_ingestion_service_failed_cleanup.py
Original file line number Diff line number Diff line change
Expand Up @@ -318,3 +318,131 @@ async def capture(*args, **kwargs):
await service.process_document("doc_10")

assert seen["result"].content == "ab"


class ExpiringUnitOfWork(PoisonedSessionUnitOfWork):
"""Like AsyncSession.rollback(): every loaded ORM object is expired, so
reading any attribute afterwards needs IO (MissingGreenlet in async code)."""

def __init__(self, repository: FakeDocumentRepositoryForFailure) -> None:
super().__init__()
self.repository = repository

async def rollback(self) -> None:
from sqlalchemy import inspect

await super().rollback()
generation = self.repository.generation
if generation is not None:
state = inspect(generation)
state._expire(state.dict, set())


@pytest.mark.asyncio
async def test_process_document_failure_handler_does_not_read_expired_generation():
document = StubDocument(
id="doc_11",
tenant_id="tenant-1",
status=DocumentStatus.INGESTED,
storage_path="tenant-1/doc_11/file.txt",
filename="file.txt",
content_hash="hash-11",
metadata_={},
)
repository = FakeDocumentRepositoryForFailure(document)
uow = ExpiringUnitOfWork(repository)
service = make_service(vector_store=FakeVectorStore(), neo4j_client=FakeNeo4jClient())
service.document_repository = repository
service.unit_of_work = uow
service.storage = PoisoningStorage(uow)

with pytest.raises(ValueError, match="storage is down"):
await service.process_document("doc_11")

assert uow.rollbacks >= 1
assert document.status == DocumentStatus.FAILED
assert document.error_message
assert document.processing_attempt_id is None


class ExpiringPoisoningStorage(PoisoningStorage):
"""A failed flush can expire ORM state before the handler starts."""

def __init__(self, uow, repository):
super().__init__(uow)
self.repository = repository

def get_file(self, storage_path):
from sqlalchemy import inspect

generation = self.repository.generation
self.generation_id = generation.id
state = inspect(generation)
state._expire(state.dict, set())
return super().get_file(storage_path)


@pytest.mark.asyncio
@pytest.mark.parametrize("existing_generation", [False, True])
@pytest.mark.parametrize("preserve_published", [False, True])
async def test_failure_handler_handles_generation_already_expired_before_rollback(
existing_generation, preserve_published,
):
document = StubDocument(
id="doc_expired_flush",
tenant_id="tenant-1",
status=DocumentStatus.READY if preserve_published else DocumentStatus.INGESTED,
storage_path="tenant-1/doc_expired_flush/file.txt",
filename="file.txt",
content_hash="hash_expired_flush",
metadata_={},
active_generation_id="published_gen" if preserve_published else None,
)
repository = FakeDocumentRepositoryForFailure(document)
if existing_generation:
repository.generation = service_module.DocumentGeneration(
id="gen_existing",
document_id=document.id,
tenant_id=document.tenant_id,
filename=document.filename,
content_hash=document.content_hash,
storage_path=document.storage_path,
metadata_={},
)
document.pending_generation_id = "gen_existing"

async def delete_chunks(generation_id):
return 0

repository.delete_chunks_by_generation = delete_chunks

failed_ids = []
mark_failed = repository.mark_generation_failed

async def record_failed(generation_id, error_message):
failed_ids.append(generation_id)
await mark_failed(generation_id, error_message)

repository.mark_generation_failed = record_failed
uow = ExpiringUnitOfWork(repository)
service = make_service(vector_store=FakeVectorStore(), neo4j_client=FakeNeo4jClient())
service.document_repository = repository
service.unit_of_work = uow
service.storage = ExpiringPoisoningStorage(uow, repository)

with pytest.raises(ValueError, match="storage is down"):
await service.process_document(document.id, force=preserve_published)

assert uow.rollbacks >= 1
assert repository.generation.status == "failed"
assert repository.generation.error_message
assert document.processing_attempt_id is None
if preserve_published:
assert document.status == DocumentStatus.READY
assert document.active_generation_id == "published_gen"
assert document.pending_generation_id is None
assert service.vector_store.delete_calls == []
else:
assert document.status == DocumentStatus.FAILED
assert document.error_message
assert failed_ids == [service.storage.generation_id]
Loading