diff --git a/api/core/rag/data_post_processor/data_post_processor.py b/api/core/rag/data_post_processor/data_post_processor.py index ace81bd79e1..56d01986510 100644 --- a/api/core/rag/data_post_processor/data_post_processor.py +++ b/api/core/rag/data_post_processor/data_post_processor.py @@ -10,6 +10,7 @@ from core.rag.rerank.entity.weight import KeywordSetting, VectorSetting, Weights from core.rag.rerank.rerank_base import BaseRerankRunner from core.rag.rerank.rerank_factory import RerankRunnerFactory from core.rag.rerank.rerank_type import RerankMode +from extensions.otel import trace_span from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.errors.invoke import InvokeAuthorizationError @@ -52,6 +53,7 @@ class DataPostProcessor: ) self.reorder_runner = self._get_reorder_runner(reorder_enabled) + @trace_span() def invoke( self, query: str, diff --git a/api/core/rag/datasource/retrieval_service.py b/api/core/rag/datasource/retrieval_service.py index ea6eee6f68a..15854a44c47 100644 --- a/api/core/rag/datasource/retrieval_service.py +++ b/api/core/rag/datasource/retrieval_service.py @@ -1,12 +1,10 @@ import concurrent.futures -import functools import logging -from collections.abc import Callable, Sequence +from collections.abc import Sequence from concurrent.futures import ThreadPoolExecutor from typing import Any, NotRequired, TypedDict from flask import Flask, current_app -from opentelemetry import context as otel_context from sqlalchemy import select from sqlalchemy.orm import Session, load_only @@ -26,7 +24,7 @@ from core.rag.rerank.rerank_type import RerankMode from core.rag.retrieval.retrieval_methods import RetrievalMethod from core.tools.signature import sign_upload_file_preview_url from extensions.ext_database import db -from extensions.otel import trace_span +from extensions.otel import propagate_context, trace_span from graphon.model_runtime.entities.model_entities import ModelType from models.dataset import ( ChildChunk, @@ -92,20 +90,6 @@ default_retrieval_model: DefaultRetrievalModelDict = { logger = logging.getLogger(__name__) -def _propagate_otel_context[**P, R](func: Callable[P, R]) -> Callable[P, R]: - captured_context = otel_context.get_current() - - @functools.wraps(func) - def wrapper(*args: P.args, **kwargs: P.kwargs) -> R: - token = otel_context.attach(captured_context) - try: - return func(*args, **kwargs) - finally: - otel_context.detach(token) - - return wrapper - - class RetrievalService: # Cache precompiled regular expressions to avoid repeated compilation @classmethod @@ -139,7 +123,7 @@ class RetrievalService: if query: futures.append( executor.submit( - _propagate_otel_context(retrieval_service._retrieve), + propagate_context(retrieval_service._retrieve), flask_app=current_app._get_current_object(), # type: ignore retrieval_method=retrieval_method, dataset=dataset, @@ -159,7 +143,7 @@ class RetrievalService: for attachment_id in attachment_ids: futures.append( executor.submit( - _propagate_otel_context(retrieval_service._retrieve), + propagate_context(retrieval_service._retrieve), flask_app=current_app._get_current_object(), # type: ignore retrieval_method=retrieval_method, dataset=dataset, @@ -820,7 +804,7 @@ class RetrievalService: if retrieval_method == RetrievalMethod.KEYWORD_SEARCH and query: futures.append( executor.submit( - _propagate_otel_context(self.keyword_search), + propagate_context(self.keyword_search), flask_app=current_app._get_current_object(), # type: ignore dataset_id=dataset.id, query=query, @@ -834,7 +818,7 @@ class RetrievalService: if query: futures.append( executor.submit( - _propagate_otel_context(self.embedding_search), + propagate_context(self.embedding_search), flask_app=current_app._get_current_object(), # type: ignore dataset_id=dataset.id, query=query, @@ -851,7 +835,7 @@ class RetrievalService: if attachment_id: futures.append( executor.submit( - _propagate_otel_context(self.embedding_search), + propagate_context(self.embedding_search), flask_app=current_app._get_current_object(), # type: ignore dataset_id=dataset.id, query=attachment_id, @@ -868,7 +852,7 @@ class RetrievalService: if RetrievalMethod.is_support_fulltext_search(retrieval_method) and query: futures.append( executor.submit( - _propagate_otel_context(self.full_text_index_search), + propagate_context(self.full_text_index_search), flask_app=current_app._get_current_object(), # type: ignore dataset_id=dataset.id, query=query, diff --git a/api/core/rag/rerank/rerank_model.py b/api/core/rag/rerank/rerank_model.py index ae7ffaad248..3dc0517860f 100644 --- a/api/core/rag/rerank/rerank_model.py +++ b/api/core/rag/rerank/rerank_model.py @@ -9,6 +9,7 @@ from core.rag.index_processor.constant.query_type import QueryType from core.rag.models.document import Document from core.rag.rerank.rerank_base import BaseRerankRunner from extensions.ext_storage import storage +from extensions.otel import trace_span from graphon.model_runtime.entities.model_entities import ModelType from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult from models.model import UploadFile @@ -22,6 +23,7 @@ class RerankModelRunner(BaseRerankRunner): self._session = session @override + @trace_span() def run( self, query: str, diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index c758bee219d..b89931f57ff 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -65,6 +65,7 @@ from core.workflow.nodes.knowledge_retrieval.retrieval import ( ) from extensions.ext_database import db from extensions.ext_redis import redis_client +from extensions.otel import propagate_context, trace_span from graphon.file import File, FileTransferMethod, FileType from graphon.model_runtime.entities.llm_entities import LLMMode, LLMResult, LLMUsage from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageRole, PromptMessageTool @@ -116,6 +117,7 @@ class DatasetRetrieval: else: self._llm_usage = self._llm_usage.plus(usage) + @trace_span() def knowledge_retrieval(self, session: Session, request: KnowledgeRetrievalRequest) -> list[Source]: self._check_knowledge_rate_limit(request.tenant_id) available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids) @@ -599,6 +601,7 @@ class DatasetRetrieval: return "\n".join([document_context.content for document_context in document_context_list]), context_files return "", context_files + @trace_span() def single_retrieve( self, session: Session, @@ -724,7 +727,7 @@ class DatasetRetrieval: if results: thread = threading.Thread( - target=self._on_retrieval_end, + target=propagate_context(self._on_retrieval_end), kwargs={ "flask_app": current_app._get_current_object(), # type: ignore "documents": results, @@ -737,6 +740,7 @@ class DatasetRetrieval: return results return [] + @trace_span() def multiple_retrieve( self, app_id: str, @@ -798,7 +802,7 @@ class DatasetRetrieval: if query: query_thread = threading.Thread( - target=self._multiple_retrieve_thread, + target=propagate_context(self._multiple_retrieve_thread_safely), kwargs={ "flask_app": current_app._get_current_object(), # type: ignore "available_datasets": available_datasets, @@ -824,7 +828,7 @@ class DatasetRetrieval: if attachment_ids: for attachment_id in attachment_ids: attachment_thread = threading.Thread( - target=self._multiple_retrieve_thread, + target=propagate_context(self._multiple_retrieve_thread_safely), kwargs={ "flask_app": current_app._get_current_object(), # type: ignore "available_datasets": available_datasets, @@ -865,7 +869,7 @@ class DatasetRetrieval: if all_documents: # add thread to call _on_retrieval_end retrieval_end_thread = threading.Thread( - target=self._on_retrieval_end, + target=propagate_context(self._on_retrieval_end), kwargs={ "flask_app": current_app._get_current_object(), # type: ignore "documents": all_documents, @@ -1161,6 +1165,7 @@ class DatasetRetrieval: all_documents.extend(documents) + @trace_span() def _run_retriever_thread( self, *, @@ -1172,27 +1177,51 @@ class DatasetRetrieval: document_ids_filter: list[str] | None, metadata_condition: MetadataFilteringCondition | None, attachment_ids: list[str] | None, + ) -> None: + with session_factory.create_session() as session: + self._retriever( + flask_app=flask_app, + session=session, + dataset_id=dataset_id, + query=query or "", + top_k=top_k, + all_documents=all_documents, + document_ids_filter=document_ids_filter, + metadata_condition=metadata_condition, + attachment_ids=attachment_ids, + ) + + def _run_retriever_thread_safely( + self, + *, + flask_app: Flask, + dataset_id: str, + query: str | None, + top_k: int, + all_documents: list[Document], + document_ids_filter: list[str] | None, + metadata_condition: MetadataFilteringCondition | None, + attachment_ids: list[str] | None, cancel_event: threading.Event | None, thread_exceptions: list[Exception] | None, ) -> None: + """Collect errors only after they pass through the traced retrieval method.""" try: - with session_factory.create_session() as session: - self._retriever( - flask_app=flask_app, - session=session, - dataset_id=dataset_id, - query=query or "", - top_k=top_k, - all_documents=all_documents, - document_ids_filter=document_ids_filter, - metadata_condition=metadata_condition, - attachment_ids=attachment_ids, - ) - except Exception as e: + self._run_retriever_thread( + flask_app=flask_app, + dataset_id=dataset_id, + query=query, + top_k=top_k, + all_documents=all_documents, + document_ids_filter=document_ids_filter, + metadata_condition=metadata_condition, + attachment_ids=attachment_ids, + ) + except Exception as exc: if cancel_event: cancel_event.set() if thread_exceptions is not None: - thread_exceptions.append(e) + thread_exceptions.append(exc) def to_dataset_retriever_tool( self, @@ -1795,6 +1824,7 @@ class DatasetRetrieval: return full_text, usage + @trace_span() def _multiple_retrieve_thread( self, flask_app: Flask, @@ -1813,11 +1843,11 @@ class DatasetRetrieval: attachment_id: str | None, dataset_count: int, cancel_event: threading.Event | None = None, - thread_exceptions: list[Exception] | None = None, - ): + ) -> None: try: with flask_app.app_context(): threads = [] + retrieval_thread_exceptions: list[Exception] = [] all_documents_item: list[Document] = [] index_type = None for dataset in available_datasets: @@ -1836,7 +1866,7 @@ class DatasetRetrieval: else: continue retrieval_thread = threading.Thread( - target=self._run_retriever_thread, + target=propagate_context(self._run_retriever_thread_safely), kwargs={ "flask_app": flask_app, "dataset_id": dataset.id, @@ -1847,7 +1877,7 @@ class DatasetRetrieval: "metadata_condition": metadata_condition, "attachment_ids": [attachment_id] if attachment_id else None, "cancel_event": cancel_event, - "thread_exceptions": thread_exceptions, + "thread_exceptions": retrieval_thread_exceptions, }, ) threads.append(retrieval_thread) @@ -1862,6 +1892,9 @@ class DatasetRetrieval: if cancel_event and cancel_event.is_set(): break + if retrieval_thread_exceptions: + raise retrieval_thread_exceptions[0] + # Skip second reranking when there is only one dataset if reranking_enable and dataset_count > 1: # do rerank for searched documents @@ -1902,11 +1935,55 @@ class DatasetRetrieval: all_documents_item = all_documents_item[:top_k] if top_k else all_documents_item if all_documents_item: all_documents.extend(all_documents_item) - except Exception as e: + except Exception: + raise + + def _multiple_retrieve_thread_safely( + self, + *, + flask_app: Flask, + available_datasets: list[Dataset], + metadata_condition: MetadataFilteringCondition | None, + metadata_filter_document_ids: dict[str, list[str]] | None, + all_documents: list[Document], + tenant_id: str, + reranking_enable: bool, + reranking_mode: str, + reranking_model: RerankingModelDict | None, + weights: WeightsDict | None, + top_k: int, + score_threshold: float, + query: str | None, + attachment_id: str | None, + dataset_count: int, + cancel_event: threading.Event | None = None, + thread_exceptions: list[Exception] | None = None, + ) -> None: + """Collect errors only after they pass through the traced multi-retrieval method.""" + try: + self._multiple_retrieve_thread( + flask_app=flask_app, + available_datasets=available_datasets, + metadata_condition=metadata_condition, + metadata_filter_document_ids=metadata_filter_document_ids, + all_documents=all_documents, + tenant_id=tenant_id, + reranking_enable=reranking_enable, + reranking_mode=reranking_mode, + reranking_model=reranking_model, + weights=weights, + top_k=top_k, + score_threshold=score_threshold, + query=query, + attachment_id=attachment_id, + dataset_count=dataset_count, + cancel_event=cancel_event, + ) + except Exception as exc: if cancel_event: cancel_event.set() if thread_exceptions is not None: - thread_exceptions.append(e) + thread_exceptions.append(exc) def _get_available_datasets(self, tenant_id: str, dataset_ids: list[str]) -> list[Dataset]: with session_factory.create_session() as session: diff --git a/api/extensions/otel/__init__.py b/api/extensions/otel/__init__.py index a431698d3d1..bb0b9b1a4b4 100644 --- a/api/extensions/otel/__init__.py +++ b/api/extensions/otel/__init__.py @@ -1,3 +1,4 @@ +from extensions.otel.context import propagate_context from extensions.otel.decorators.base import trace_span from extensions.otel.decorators.handler import SpanHandler from extensions.otel.decorators.handlers.generate_handler import AppGenerateHandler @@ -7,5 +8,6 @@ __all__ = [ "AppGenerateHandler", "SpanHandler", "WorkflowAppRunnerHandler", + "propagate_context", "trace_span", ] diff --git a/api/extensions/otel/context.py b/api/extensions/otel/context.py new file mode 100644 index 00000000000..b7378a3e667 --- /dev/null +++ b/api/extensions/otel/context.py @@ -0,0 +1,21 @@ +"""Utilities for propagating OpenTelemetry context across execution boundaries.""" + +import functools +from collections.abc import Callable + +from opentelemetry import context as otel_context + + +def propagate_context[**P, R](func: Callable[P, R]) -> Callable[P, R]: + """Capture the current context and attach it whenever ``func`` executes.""" + captured_context = otel_context.get_current() + + @functools.wraps(func) + def wrapper(*args: P.args, **kwargs: P.kwargs) -> R: + token = otel_context.attach(captured_context) + try: + return func(*args, **kwargs) + finally: + otel_context.detach(token) + + return wrapper diff --git a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py index bebd9face61..9e504dd1b9b 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_dataset_retrieval.py @@ -3752,8 +3752,8 @@ class TestKnowledgeRetrievalRegression: """ Repro test for current bug: reranking runs after `with flask_app.app_context():` exits. - `_multiple_retrieve_thread` catches exceptions and stores them into `thread_exceptions`, - so we must assert from that list (not from an outer try/except). + The outer thread entry point catches exceptions from the traced retrieval method + and stores them in `thread_exceptions`. """ dataset_retrieval = DatasetRetrieval() flask_app = Flask(__name__) @@ -3806,7 +3806,6 @@ class TestKnowledgeRetrievalRegression: # output list from _multiple_retrieve_thread all_documents: list[Document] = [] - # IMPORTANT: _multiple_retrieve_thread swallows exceptions and appends them here thread_exceptions: list[Exception] = [] def target(): @@ -3818,7 +3817,7 @@ class TestKnowledgeRetrievalRegression: ), _patched_retriever_session(), ): - dataset_retrieval._multiple_retrieve_thread( + dataset_retrieval._multiple_retrieve_thread_safely( flask_app=flask_app, available_datasets=[mock_dataset, secondary_dataset], metadata_condition=None, @@ -3847,7 +3846,6 @@ class TestKnowledgeRetrievalRegression: # Ensure reranking branch was actually executed assert called["init"] >= 1, "DataPostProcessor was never constructed; reranking branch may not have run." - # Current buggy code should record an exception (not raise it) assert not thread_exceptions, thread_exceptions def test_run_retriever_thread_provides_session_to_retriever(self): @@ -3865,14 +3863,12 @@ class TestKnowledgeRetrievalRegression: document_ids_filter=None, metadata_condition=None, attachment_ids=None, - cancel_event=None, - thread_exceptions=[], ) mock_retriever.assert_called_once() assert mock_retriever.call_args.kwargs["session"] is session - def test_run_retriever_thread_records_retriever_exception(self): + def test_run_retriever_thread_safely_records_retriever_exception(self): dataset_retrieval = DatasetRetrieval() all_documents: list[Document] = [] cancel_event = threading.Event() @@ -3881,7 +3877,7 @@ class TestKnowledgeRetrievalRegression: with _patched_retriever_session(): with patch.object(dataset_retrieval, "_retriever", side_effect=expected_error): - dataset_retrieval._run_retriever_thread( + dataset_retrieval._run_retriever_thread_safely( flask_app=_FakeFlaskApp(), dataset_id="dataset-1", query="test query", @@ -5139,7 +5135,7 @@ class TestSingleAndMultipleRetrieveCoverage: app = Flask(__name__) def failing_thread(**kwargs): - kwargs["thread_exceptions"].append(RuntimeError("thread boom")) + raise RuntimeError("thread boom") with app.app_context(): with ( diff --git a/api/tests/unit_tests/extensions/otel/test_context.py b/api/tests/unit_tests/extensions/otel/test_context.py new file mode 100644 index 00000000000..7a837d13ee1 --- /dev/null +++ b/api/tests/unit_tests/extensions/otel/test_context.py @@ -0,0 +1,51 @@ +from concurrent.futures import ThreadPoolExecutor + +import pytest +from opentelemetry import context as otel_context + +from extensions.otel.context import propagate_context + + +def test_propagate_context_captures_context_when_wrapped() -> None: + context_key = otel_context.create_key("test-context") + captured_context = otel_context.set_value(context_key, "captured") + + token = otel_context.attach(captured_context) + try: + wrapped = propagate_context(lambda: otel_context.get_value(context_key)) + finally: + otel_context.detach(token) + + with ThreadPoolExecutor(max_workers=1) as executor: + assert executor.submit(wrapped).result() == "captured" + + +def test_propagate_context_detaches_context_after_exception() -> None: + context_key = otel_context.create_key("test-context") + captured_context = otel_context.set_value(context_key, "captured") + + def raise_error() -> None: + raise RuntimeError("retrieval failed") + + token = otel_context.attach(captured_context) + try: + wrapped = propagate_context(raise_error) + finally: + otel_context.detach(token) + + def invoke_and_read_context() -> str | None: + with pytest.raises(RuntimeError, match="retrieval failed"): + wrapped() + return otel_context.get_value(context_key) + + with ThreadPoolExecutor(max_workers=1) as executor: + assert executor.submit(invoke_and_read_context).result() is None + + +def test_propagate_context_preserves_function_metadata() -> None: + def retrieve() -> None: + pass + + wrapped = propagate_context(retrieve) + + assert wrapped.__name__ == "retrieve" diff --git a/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py b/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py new file mode 100644 index 00000000000..586d39f8f60 --- /dev/null +++ b/api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py @@ -0,0 +1,122 @@ +import threading +from unittest.mock import MagicMock, patch +from uuid import uuid4 + +from opentelemetry.trace import StatusCode, get_current_span, get_tracer + +from core.rag.rerank.rerank_type import RerankMode +from core.rag.retrieval.dataset_retrieval import DatasetRetrieval +from core.workflow.nodes.knowledge_retrieval.retrieval import KnowledgeRetrievalRequest +from models.dataset import Dataset + + +def test_knowledge_retrieval_creates_a_child_otel_span( + memory_span_exporter, + tracer_provider_with_memory_exporter, +) -> None: + """The retrieval entry point must be visible beneath its workflow node span.""" + request = KnowledgeRetrievalRequest( + tenant_id=str(uuid4()), + user_id=str(uuid4()), + app_id=str(uuid4()), + user_from="account", + dataset_ids=[str(uuid4())], + retrieval_mode="multiple", + query="test query", + ) + retrieval = DatasetRetrieval() + + with ( + patch("extensions.otel.decorators.base.dify_config.ENABLE_OTEL", True), + patch.object(retrieval, "_check_knowledge_rate_limit"), + patch.object(retrieval, "_get_available_datasets", return_value=[]), + get_tracer(__name__).start_as_current_span("knowledge-retrieval-node") as node_span, + ): + assert retrieval.knowledge_retrieval(MagicMock(), request) == [] + + retrieval_span = next( + span + for span in memory_span_exporter.get_finished_spans() + if span.name == "core.rag.retrieval.dataset_retrieval.DatasetRetrieval.knowledge_retrieval" + ) + node_span_context = node_span.get_span_context() + assert retrieval_span.context.trace_id == node_span_context.trace_id + assert retrieval_span.parent is not None + assert retrieval_span.parent.span_id == node_span_context.span_id + + +def test_multiple_retrieve_preserves_otel_context_in_dataset_thread( + app, + tracer_provider_with_memory_exporter, +) -> None: + """Per-dataset retrieval spans must remain in the workflow node trace.""" + retrieval = DatasetRetrieval() + dataset = MagicMock(spec=Dataset) + dataset.id = str(uuid4()) + dataset.indexing_technique = "high_quality" + dataset.embedding_model = "text-embedding-3-small" + dataset.embedding_model_provider = "openai" + observed_trace_ids: list[int] = [] + + def record_active_trace(**_kwargs: object) -> None: + observed_trace_ids.append(get_current_span().get_span_context().trace_id) + + with ( + app.app_context(), + patch("extensions.otel.decorators.base.dify_config.ENABLE_OTEL", True), + patch.object(retrieval, "_multiple_retrieve_thread", side_effect=record_active_trace), + patch.object(retrieval, "_on_query"), + get_tracer(__name__).start_as_current_span("knowledge-retrieval-node") as node_span, + ): + retrieval.multiple_retrieve( + app_id=str(uuid4()), + tenant_id=str(uuid4()), + user_id=str(uuid4()), + user_from="account", + available_datasets=[dataset], + query="test query", + top_k=4, + score_threshold=0.0, + reranking_mode=RerankMode.RERANKING_MODEL, + reranking_enable=False, + ) + + assert observed_trace_ids == [node_span.get_span_context().trace_id] + + +def test_retriever_thread_exception_sets_error_span_and_is_collected( + 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") + + 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), + ): + retrieval._run_retriever_thread_safely( + flask_app=app, + dataset_id=str(uuid4()), + 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, + ) + + retrieval_span = next( + span + for span in memory_span_exporter.get_finished_spans() + if span.name.endswith("DatasetRetrieval._run_retriever_thread") + ) + assert retrieval_span.status.status_code == StatusCode.ERROR + assert cancel_event.is_set() + assert thread_exceptions == [expected_error]