mirror of
https://github.com/langgenius/dify.git
synced 2026-09-01 15:09:21 +08:00
fix: _multiple_retrieve_thread not fail fast (#40466)
This commit is contained in:
@@ -9,6 +9,7 @@ from collections.abc import Generator, Mapping
|
||||
from typing import Any, Union, cast
|
||||
|
||||
from flask import Flask, current_app
|
||||
from opentelemetry.trace import get_current_span
|
||||
from sqlalchemy import and_, func, literal, or_, select, update
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
@@ -1204,8 +1205,9 @@ class DatasetRetrieval:
|
||||
attachment_ids: list[str] | None,
|
||||
cancel_event: threading.Event | None,
|
||||
thread_exceptions: list[Exception] | None,
|
||||
skip_on_error: bool = False,
|
||||
) -> None:
|
||||
"""Collect errors only after they pass through the traced retrieval method."""
|
||||
"""Collect errors after tracing, or skip dataset-level failures when requested."""
|
||||
try:
|
||||
self._run_retriever_thread(
|
||||
flask_app=flask_app,
|
||||
@@ -1218,6 +1220,25 @@ class DatasetRetrieval:
|
||||
attachment_ids=attachment_ids,
|
||||
)
|
||||
except Exception as exc:
|
||||
if skip_on_error:
|
||||
logger.warning(
|
||||
"Skipping dataset retrieval because retriever failed, dataset_id=%s, error_type=%s, error=%s",
|
||||
dataset_id,
|
||||
type(exc).__name__,
|
||||
str(exc),
|
||||
)
|
||||
span = get_current_span()
|
||||
if span and span.is_recording():
|
||||
span.add_event(
|
||||
"dataset_retrieval.skipped",
|
||||
attributes={
|
||||
"dataset_id": dataset_id,
|
||||
"error.type": type(exc).__name__,
|
||||
"error.message": str(exc),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if cancel_event:
|
||||
cancel_event.set()
|
||||
if thread_exceptions is not None:
|
||||
@@ -1878,6 +1899,7 @@ class DatasetRetrieval:
|
||||
"attachment_ids": [attachment_id] if attachment_id else None,
|
||||
"cancel_event": cancel_event,
|
||||
"thread_exceptions": retrieval_thread_exceptions,
|
||||
"skip_on_error": True,
|
||||
},
|
||||
)
|
||||
threads.append(retrieval_thread)
|
||||
|
||||
@@ -3894,6 +3894,96 @@ class TestKnowledgeRetrievalRegression:
|
||||
assert cancel_event.is_set()
|
||||
assert thread_exceptions == [expected_error]
|
||||
|
||||
def test_run_retriever_thread_safely_skips_failed_dataset_when_requested(self, caplog):
|
||||
dataset_retrieval = DatasetRetrieval()
|
||||
all_documents: list[Document] = []
|
||||
cancel_event = threading.Event()
|
||||
thread_exceptions: list[Exception] = []
|
||||
expected_error = RuntimeError("retrieval failed")
|
||||
|
||||
with _patched_retriever_session():
|
||||
with patch.object(dataset_retrieval, "_retriever", side_effect=expected_error):
|
||||
dataset_retrieval._run_retriever_thread_safely(
|
||||
flask_app=_FakeFlaskApp(),
|
||||
dataset_id="dataset-1",
|
||||
query="test query",
|
||||
top_k=3,
|
||||
all_documents=all_documents,
|
||||
document_ids_filter=None,
|
||||
metadata_condition=None,
|
||||
attachment_ids=None,
|
||||
cancel_event=cancel_event,
|
||||
thread_exceptions=thread_exceptions,
|
||||
skip_on_error=True,
|
||||
)
|
||||
|
||||
assert not cancel_event.is_set()
|
||||
assert thread_exceptions == []
|
||||
assert "dataset_id=dataset-1" in caplog.text
|
||||
assert "Skipping dataset retrieval because retriever failed" in caplog.text
|
||||
|
||||
def test_multiple_retrieve_thread_skips_failed_dataset(self, mock_dataset, caplog):
|
||||
dataset_retrieval = DatasetRetrieval()
|
||||
flask_app = Flask(__name__)
|
||||
successful_dataset = Dataset(
|
||||
id=str(uuid4()),
|
||||
provider="dify",
|
||||
indexing_technique="high_quality",
|
||||
)
|
||||
document = Document(
|
||||
page_content="successful doc",
|
||||
metadata={
|
||||
"doc_id": "doc1",
|
||||
"score": 0.95,
|
||||
"document_id": str(uuid4()),
|
||||
"dataset_id": successful_dataset.id,
|
||||
},
|
||||
provider="dify",
|
||||
)
|
||||
|
||||
def fake_retriever(
|
||||
flask_app,
|
||||
session,
|
||||
dataset_id,
|
||||
query,
|
||||
top_k,
|
||||
all_documents,
|
||||
document_ids_filter,
|
||||
metadata_condition,
|
||||
attachment_ids,
|
||||
):
|
||||
if dataset_id == mock_dataset.id:
|
||||
raise RuntimeError("dataset unavailable")
|
||||
all_documents.append(document)
|
||||
|
||||
all_documents: list[Document] = []
|
||||
|
||||
with (
|
||||
patch.object(dataset_retrieval, "_retriever", side_effect=fake_retriever),
|
||||
_patched_retriever_session(),
|
||||
):
|
||||
dataset_retrieval._multiple_retrieve_thread(
|
||||
flask_app=flask_app,
|
||||
available_datasets=[mock_dataset, successful_dataset],
|
||||
metadata_condition=None,
|
||||
metadata_filter_document_ids=None,
|
||||
all_documents=all_documents,
|
||||
tenant_id=str(uuid4()),
|
||||
reranking_enable=False,
|
||||
reranking_mode="reranking_model",
|
||||
reranking_model=None,
|
||||
weights=None,
|
||||
top_k=3,
|
||||
score_threshold=0.0,
|
||||
query="test query",
|
||||
attachment_id=None,
|
||||
dataset_count=2,
|
||||
)
|
||||
|
||||
assert all_documents == [document]
|
||||
assert f"dataset_id={mock_dataset.id}" in caplog.text
|
||||
assert "Skipping dataset retrieval because retriever failed" in caplog.text
|
||||
|
||||
|
||||
class _FakeFlaskApp:
|
||||
def app_context(self):
|
||||
|
||||
@@ -121,3 +121,47 @@ def test_retriever_thread_exception_sets_error_span_and_is_collected(
|
||||
assert retrieval_span.status.status_code == StatusCode.ERROR
|
||||
assert cancel_event.is_set()
|
||||
assert thread_exceptions == [expected_error]
|
||||
|
||||
|
||||
def test_retriever_thread_exception_emits_skip_event_when_requested(
|
||||
app,
|
||||
memory_span_exporter,
|
||||
tracer_provider_with_memory_exporter,
|
||||
) -> None:
|
||||
retrieval = DatasetRetrieval()
|
||||
cancel_event = threading.Event()
|
||||
thread_exceptions: list[Exception] = []
|
||||
expected_error = RuntimeError("retrieval failed")
|
||||
dataset_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("extensions.otel.decorators.base.dify_config.ENABLE_OTEL", True),
|
||||
patch("core.rag.retrieval.dataset_retrieval.session_factory.create_session"),
|
||||
patch.object(retrieval, "_retriever", side_effect=expected_error),
|
||||
get_tracer(__name__).start_as_current_span("dataset-retrieval-parent") as parent_span,
|
||||
):
|
||||
retrieval._run_retriever_thread_safely(
|
||||
flask_app=app,
|
||||
dataset_id=dataset_id,
|
||||
query="test query",
|
||||
top_k=4,
|
||||
all_documents=[],
|
||||
document_ids_filter=None,
|
||||
metadata_condition=None,
|
||||
attachment_ids=None,
|
||||
cancel_event=cancel_event,
|
||||
thread_exceptions=thread_exceptions,
|
||||
skip_on_error=True,
|
||||
)
|
||||
|
||||
retrieval_span = next(
|
||||
span
|
||||
for span in memory_span_exporter.get_finished_spans()
|
||||
if span.name.endswith("DatasetRetrieval._run_retriever_thread")
|
||||
)
|
||||
skip_event = next(event for event in parent_span.events if event.name == "dataset_retrieval.skipped")
|
||||
assert retrieval_span.status.status_code == StatusCode.ERROR
|
||||
assert skip_event.attributes["dataset_id"] == dataset_id
|
||||
assert skip_event.attributes["error.message"] == "retrieval failed"
|
||||
assert not cancel_event.is_set()
|
||||
assert thread_exceptions == []
|
||||
|
||||
Reference in New Issue
Block a user