mirror of
https://github.com/langgenius/dify.git
synced 2026-07-25 21:48:30 +08:00
feat: knowledge add more trace (#38959)
This commit is contained in:
parent
14c50a9f09
commit
db29caff2b
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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,
|
||||
|
||||
@ -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:
|
||||
|
||||
@ -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",
|
||||
]
|
||||
|
||||
21
api/extensions/otel/context.py
Normal file
21
api/extensions/otel/context.py
Normal file
@ -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
|
||||
@ -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 (
|
||||
|
||||
51
api/tests/unit_tests/extensions/otel/test_context.py
Normal file
51
api/tests/unit_tests/extensions/otel/test_context.py
Normal file
@ -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"
|
||||
122
api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py
Normal file
122
api/tests/unit_tests/extensions/otel/test_retrieval_tracing.py
Normal file
@ -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]
|
||||
Loading…
Reference in New Issue
Block a user