fix: make document retry batches consistent (#39978)

Co-authored-by: caosurong <surong.cao@thinkingdata.cn>
This commit is contained in:
csurong
2026-08-06 02:01:14 +00:00
committed by GitHub
co-authored by caosurong
parent b6b3e747f3
commit f7881bb635
5 changed files with 251 additions and 33 deletions
@@ -1441,9 +1441,11 @@ class DocumentRetryApi(DocumentResource):
retry_documents = []
if not dataset:
raise NotFound("Dataset not found.")
documents = DocumentService.get_documents_by_ids(dataset.id, payload.document_ids, session)
documents_by_id = {document.id: document for document in documents}
for document_id in payload.document_ids:
try:
document = DocumentService.get_document(dataset.id, document_id, session=session)
document = documents_by_id.get(document_id)
# 404 if document not found
if document is None:
+37 -13
View File
@@ -2099,22 +2099,46 @@ class DocumentService:
@staticmethod
def retry_document(dataset_id: str, documents: list[Document], session: Session):
for document in documents:
# add retry flag
retry_indexing_cache_key = f"document_{document.id}_is_retried"
cache_result = redis_client.get(retry_indexing_cache_key)
if cache_result is not None:
raise ValueError("Document is being retried, please try again later")
# retry document indexing
document.indexing_status = IndexingStatus.WAITING
session.add(document)
session.commit()
"""Reserve the whole retry batch before changing any document state.
redis_client.setex(retry_indexing_cache_key, 600, 1)
# trigger async task
document_ids = [document.id for document in documents]
Redis lock acquisition is intentionally coupled to this bounded status
transaction so a concurrent request cannot partially admit the batch.
"""
if not current_user or not current_user.id:
raise ValueError("Current user or current user id not found")
unique_documents = list({document.id: document for document in documents}.values())
retry_indexing_cache_keys = [f"document_{document.id}_is_retried" for document in unique_documents]
acquired_locks: list[Any] = []
def release_acquired_locks() -> None:
for retry_lock in acquired_locks:
try:
retry_lock.release()
except Exception:
logger.warning("Failed to release document retry lock", exc_info=True)
try:
for retry_indexing_cache_key in retry_indexing_cache_keys:
retry_lock = redis_client.lock(retry_indexing_cache_key, timeout=600, thread_local=False)
if not retry_lock.acquire(blocking=False):
raise ValueError("Document is being retried, please try again later")
acquired_locks.append(retry_lock)
except Exception:
release_acquired_locks()
raise
try:
for document in unique_documents:
document.indexing_status = IndexingStatus.WAITING
session.add(document)
session.commit()
except Exception:
session.rollback()
release_acquired_locks()
raise
document_ids = [document.id for document in unique_documents]
retry_document_indexing_task.delay(dataset_id, document_ids, current_user.id)
@staticmethod
@@ -8,7 +8,7 @@ pause, recover, retry, batch updates, and renaming.
import datetime
import json
from unittest.mock import create_autospec, patch
from unittest.mock import MagicMock, create_autospec, patch
from uuid import uuid4
import pytest
@@ -562,9 +562,9 @@ class TestDocumentServiceRetryDocument:
The retry_document method:
1. Validates documents are not already being retried
2. Sets retry flag in Redis cache
3. Resets document indexing_status to waiting
4. Commits changes to database
2. Atomically reserves retry flags in Redis cache
3. Resets all document indexing statuses to waiting
4. Commits all changes together
5. Triggers retry task
Test scenarios include:
@@ -595,12 +595,22 @@ class TestDocumentServiceRetryDocument:
):
user_id = str(uuid4())
mock_current_user.id = user_id
retry_locks = []
def create_retry_lock(*_args, **_kwargs):
retry_lock = MagicMock()
retry_lock.acquire.return_value = True
retry_locks.append(retry_lock)
return retry_lock
mock_redis.lock.side_effect = create_retry_lock
yield {
"current_user": mock_current_user,
"redis_client": mock_redis,
"retry_task": mock_task,
"user_id": user_id,
"retry_locks": retry_locks,
}
def test_retry_document_single_success(
@@ -629,8 +639,6 @@ class TestDocumentServiceRetryDocument:
indexing_status=IndexingStatus.ERROR,
)
mock_document_service_dependencies["redis_client"].get.return_value = None
# Act
DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers)
@@ -639,7 +647,11 @@ class TestDocumentServiceRetryDocument:
assert document.indexing_status == IndexingStatus.WAITING
expected_cache_key = f"document_{document.id}_is_retried"
mock_document_service_dependencies["redis_client"].setex.assert_called_once_with(expected_cache_key, 600, 1)
mock_document_service_dependencies["redis_client"].lock.assert_called_once_with(
expected_cache_key, timeout=600, thread_local=False
)
retry_lock = mock_document_service_dependencies["retry_locks"][0]
retry_lock.acquire.assert_called_once_with(blocking=False)
mock_document_service_dependencies["retry_task"].delay.assert_called_once_with(
dataset.id, [document.id], mock_document_service_dependencies["user_id"]
)
@@ -676,8 +688,6 @@ class TestDocumentServiceRetryDocument:
position=2,
)
mock_document_service_dependencies["redis_client"].get.return_value = None
# Act
DocumentService.retry_document(dataset.id, [document1, document2], session=db_session_with_containers)
@@ -715,7 +725,10 @@ class TestDocumentServiceRetryDocument:
indexing_status=IndexingStatus.ERROR,
)
mock_document_service_dependencies["redis_client"].get.return_value = "1"
retry_lock = MagicMock()
retry_lock.acquire.return_value = False
mock_document_service_dependencies["redis_client"].lock.side_effect = None
mock_document_service_dependencies["redis_client"].lock.return_value = retry_lock
# Act & Assert
with pytest.raises(ValueError, match="Document is being retried, please try again later"):
@@ -724,6 +737,42 @@ class TestDocumentServiceRetryDocument:
db_session_with_containers.refresh(document)
assert document.indexing_status == IndexingStatus.ERROR
def test_retry_document_later_conflict_leaves_batch_unchanged(
self, db_session_with_containers: Session, mock_document_service_dependencies
):
dataset = DocumentStatusTestDataFactory.create_dataset(db_session_with_containers)
document1 = DocumentStatusTestDataFactory.create_document(
db_session_with_containers,
dataset_id=dataset.id,
tenant_id=dataset.tenant_id,
document_id=str(uuid4()),
indexing_status=IndexingStatus.ERROR,
)
document2 = DocumentStatusTestDataFactory.create_document(
db_session_with_containers,
dataset_id=dataset.id,
tenant_id=dataset.tenant_id,
document_id=str(uuid4()),
indexing_status=IndexingStatus.ERROR,
position=2,
)
first_retry_lock = MagicMock()
first_retry_lock.acquire.return_value = True
second_retry_lock = MagicMock()
second_retry_lock.acquire.return_value = False
mock_document_service_dependencies["redis_client"].lock.side_effect = [first_retry_lock, second_retry_lock]
with pytest.raises(ValueError, match="Document is being retried, please try again later"):
DocumentService.retry_document(dataset.id, [document1, document2], session=db_session_with_containers)
db_session_with_containers.refresh(document1)
db_session_with_containers.refresh(document2)
assert document1.indexing_status == IndexingStatus.ERROR
assert document2.indexing_status == IndexingStatus.ERROR
first_retry_lock.release.assert_called_once_with()
second_retry_lock.release.assert_not_called()
mock_document_service_dependencies["retry_task"].delay.assert_not_called()
def test_retry_document_missing_current_user_error(
self, db_session_with_containers: Session, mock_document_service_dependencies
):
@@ -748,13 +797,16 @@ class TestDocumentServiceRetryDocument:
indexing_status=IndexingStatus.ERROR,
)
mock_document_service_dependencies["redis_client"].get.return_value = None
mock_document_service_dependencies["current_user"].id = None
# Act & Assert
with pytest.raises(ValueError, match="Current user or current user id not found"):
DocumentService.retry_document(dataset.id, [document], session=db_session_with_containers)
db_session_with_containers.refresh(document)
assert document.indexing_status == IndexingStatus.ERROR
mock_document_service_dependencies["redis_client"].lock.assert_not_called()
class TestDocumentServiceBatchUpdateDocumentStatus:
"""
@@ -732,11 +732,14 @@ class TestDocumentRetryApi:
api = DocumentRetryApi()
method = unwrap(api.post)
payload = {"document_ids": ["doc-1"]}
doc = MagicMock(indexing_status="indexing")
doc = MagicMock(id="doc-1", indexing_status="indexing")
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=doc),
patch(
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
return_value=[doc],
),
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=True),
patch("controllers.console.datasets.datasets_document.DocumentService.retry_document") as retry_mock,
):
@@ -748,11 +751,14 @@ class TestDocumentRetryApi:
api = DocumentRetryApi()
method = unwrap(api.post)
payload = {"document_ids": ["doc-1"]}
document = MagicMock(indexing_status=IndexingStatus.INDEXING, archived=False)
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.INDEXING, archived=False)
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document),
patch(
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
return_value=[document],
),
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
patch(
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
@@ -762,15 +768,44 @@ class TestDocumentRetryApi:
assert status == 204
retry_mock.assert_called_once_with("ds-1", [document], ANY)
def test_retry_loads_selected_documents_in_one_batch(self, app: Flask, patch_tenant, patch_dataset):
api = DocumentRetryApi()
method = unwrap(api.post)
payload = {"document_ids": ["doc-1", "doc-2"]}
first_document = MagicMock(id="doc-1", indexing_status=IndexingStatus.ERROR, archived=False)
second_document = MagicMock(id="doc-2", indexing_status=IndexingStatus.ERROR, archived=False)
session = MagicMock()
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch(
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
return_value=[first_document, second_document],
) as get_documents_by_ids,
patch("controllers.console.datasets.datasets_document.DocumentService.check_archived", return_value=False),
patch(
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
) as retry_mock,
):
response, status = method(api, session, "ds-1")
assert status == 204
get_documents_by_ids.assert_called_once_with("ds-1", ["doc-1", "doc-2"], session)
retry_mock.assert_called_once_with("ds-1", [first_document, second_document], session)
def test_retry_skips_completed_document(self, app: Flask, patch_tenant, patch_dataset):
api = DocumentRetryApi()
method = unwrap(api.post)
payload = {"document_ids": ["doc-1"]}
document = MagicMock(indexing_status=IndexingStatus.COMPLETED, archived=False)
document = MagicMock(id="doc-1", indexing_status=IndexingStatus.COMPLETED, archived=False)
with (
app.test_request_context("/", json=payload),
patch.object(type(console_ns), "payload", payload),
patch("controllers.console.datasets.datasets_document.DocumentService.get_document", return_value=document),
patch(
"controllers.console.datasets.datasets_document.DocumentService.get_documents_by_ids",
return_value=[document],
),
patch(
"controllers.console.datasets.datasets_document.DocumentService.retry_document", return_value=None
) as retry_mock,
@@ -44,6 +44,44 @@ from .dataset_service_test_helpers import (
)
class _RetryFlagLock:
def __init__(self, store: "_RetryFlagStore", key: str):
self.store = store
self.key = key
self.token = f"owner-{store.next_token}"
store.next_token += 1
def acquire(self, *, blocking: bool):
assert blocking is False
if self.key in self.store.values:
if self.store.replacement_on_conflict:
replacement_key, replacement_value = self.store.replacement_on_conflict
self.store.values[replacement_key] = replacement_value
return False
self.store.values[self.key] = self.token
return True
def release(self):
if self.store.values.get(self.key) == self.token:
self.store.values.pop(self.key)
class _RetryFlagStore:
def __init__(
self,
values: dict[str, str] | None = None,
replacement_on_conflict: tuple[str, str] | None = None,
):
self.values = values or {}
self.replacement_on_conflict = replacement_on_conflict
self.next_token = 1
def lock(self, key: str, *, timeout: int, thread_local: bool):
assert timeout == 600
assert thread_local is False
return _RetryFlagLock(self, key)
class TestDocumentServiceDisplayStatus:
"""Unit tests for DocumentService display-status helpers."""
@@ -214,16 +252,83 @@ class TestDocumentServiceMutations:
with pytest.raises(DocumentIndexingError):
DocumentService.recover_document(document, session)
def test_retry_document_raises_when_retry_flag_is_already_set(self):
def test_retry_document_raises_when_retry_flag_is_already_set(self, rename_account_context):
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1")
session = MagicMock()
with patch("services.dataset_service.redis_client") as mock_redis:
mock_redis.get.return_value = "1"
mock_redis.lock.return_value.acquire.return_value = False
with pytest.raises(ValueError, match="being retried"):
DocumentService.retry_document("dataset-1", [document], session)
def test_retry_document_leaves_batch_unchanged_when_later_document_is_already_being_retried(
self, rename_account_context
):
first_document = DatasetServiceUnitDataFactory.create_document_mock(
document_id="doc-1", indexing_status="error"
)
second_document = DatasetServiceUnitDataFactory.create_document_mock(
document_id="doc-2", indexing_status="error"
)
first_retry_key = "document_doc-1_is_retried"
second_retry_key = "document_doc-2_is_retried"
retry_flags = _RetryFlagStore({second_retry_key: "other-request"})
session = MagicMock()
with (
patch("services.dataset_service.redis_client", retry_flags),
patch("services.dataset_service.retry_document_indexing_task") as retry_task,
):
with pytest.raises(ValueError, match="being retried"):
DocumentService.retry_document(
"dataset-1",
[first_document, second_document],
session,
)
assert first_document.indexing_status == "error"
assert second_document.indexing_status == "error"
assert first_retry_key not in retry_flags.values
assert retry_flags.values[second_retry_key] == "other-request"
retry_task.delay.assert_not_called()
def test_retry_document_does_not_release_a_retry_flag_reacquired_by_another_request(self, rename_account_context):
first_retry_key = "document_doc-1_is_retried"
second_retry_key = "document_doc-2_is_retried"
retry_flags = _RetryFlagStore(
{second_retry_key: "other-request"},
replacement_on_conflict=(first_retry_key, "new-owner"),
)
documents = [
DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", indexing_status="error"),
DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-2", indexing_status="error"),
]
with patch("services.dataset_service.redis_client", retry_flags):
with pytest.raises(ValueError, match="being retried"):
DocumentService.retry_document("dataset-1", documents, MagicMock())
assert retry_flags.values[first_retry_key] == "new-owner"
assert retry_flags.values[second_retry_key] == "other-request"
def test_retry_document_releases_flags_when_status_commit_fails(self, rename_account_context):
retry_flags = _RetryFlagStore()
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1", indexing_status="error")
session = MagicMock()
session.commit.side_effect = RuntimeError("database unavailable")
with (
patch("services.dataset_service.redis_client", retry_flags),
patch("services.dataset_service.retry_document_indexing_task") as retry_task,
):
with pytest.raises(RuntimeError, match="database unavailable"):
DocumentService.retry_document("dataset-1", [document], session)
assert retry_flags.values == {}
session.rollback.assert_called_once_with()
retry_task.delay.assert_not_called()
def test_sync_website_document_raises_when_sync_flag_exists(self):
document = DatasetServiceUnitDataFactory.create_document_mock(document_id="doc-1")
session = MagicMock()