feat: knowledge add more trace (#38959)

This commit is contained in:
wangxiaolei 2026-07-24 10:27:16 +08:00 committed by GitHub
parent 14c50a9f09
commit db29caff2b
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
9 changed files with 315 additions and 58 deletions

View File

@ -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,

View File

@ -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,

View File

@ -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,

View File

@ -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:

View File

@ -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",
]

View 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

View File

@ -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 (

View 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"

View 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]