From 4a6535c014c03acf5d3981c40f3536977048cf51 Mon Sep 17 00:00:00 2001 From: Alexander Assmann Date: Mon, 31 Aug 2026 14:53:51 +0200 Subject: [PATCH] feat: add Langfuse tracing for R2R model calls --- py/core/main/api/v3/documents_router.py | 3 + .../hatchet/ingestion_workflow.py | 36 +++- .../simple/ingestion_workflow.py | 15 ++ py/core/main/services/ingestion_service.py | 113 ++++++---- py/core/main/services/retrieval_service.py | 102 +++++---- py/core/parsers/media/audio_parser.py | 18 +- py/core/parsers/media/img_parser.py | 4 +- py/core/parsers/media/pdf_parser.py | 8 + py/core/providers/embeddings/litellm.py | 6 + py/core/providers/llm/litellm.py | 11 + py/core/utils/observability.py | 204 ++++++++++++++++++ py/tests/unit/test_litellm_observability.py | 62 ++++++ py/tests/unit/test_observability.py | 115 ++++++++++ 13 files changed, 595 insertions(+), 102 deletions(-) create mode 100644 py/core/utils/observability.py create mode 100644 py/tests/unit/test_litellm_observability.py create mode 100644 py/tests/unit/test_observability.py diff --git a/py/core/main/api/v3/documents_router.py b/py/core/main/api/v3/documents_router.py index 68c7b81fa5..b6bfa4db0b 100644 --- a/py/core/main/api/v3/documents_router.py +++ b/py/core/main/api/v3/documents_router.py @@ -39,6 +39,7 @@ WrappedRelationshipsResponse, ) from core.utils import update_settings_from_dict +from core.utils.observability import create_r2r_trace_id from shared.abstractions import IngestionMode from ...abstractions import R2RProviders, R2RServices @@ -398,6 +399,7 @@ async def create_document( # Prepare workflow input workflow_input = { + "langfuse_trace_id": create_r2r_trace_id(), "document_id": str(document_id), "chunks": [ chunk.model_dump(mode="json") @@ -508,6 +510,7 @@ async def create_document( ) workflow_input = { + "langfuse_trace_id": create_r2r_trace_id(), "file_data": file_data, "document_id": str(document_id), "collection_ids": ( diff --git a/py/core/main/orchestration/hatchet/ingestion_workflow.py b/py/core/main/orchestration/hatchet/ingestion_workflow.py index 85f0311990..a1488b9ee4 100644 --- a/py/core/main/orchestration/hatchet/ingestion_workflow.py +++ b/py/core/main/orchestration/hatchet/ingestion_workflow.py @@ -22,6 +22,12 @@ num_tokens, update_settings_from_dict, ) +from core.utils.observability import ( + build_r2r_ingestion_trace_context, + r2r_ingestion_trace_context, + reset_r2r_trace_context, + set_r2r_trace_context, +) from ...services import IngestionService, IngestionServiceAdapter @@ -59,9 +65,15 @@ def concurrency(self, context: Context) -> str: @orchestration_provider.step(retries=0, timeout="60m") async def parse(self, context: Context) -> dict: + input_data = context.workflow_input()["request"] + trace_token = set_r2r_trace_context( + build_r2r_ingestion_trace_context( + input_data, + task_id=str(context.workflow_run_id()), + ) + ) try: logger.info("Initiating ingestion workflow, step: parse") - input_data = context.workflow_input()["request"] parsed_data = IngestionServiceAdapter.parse_ingest_file_input( input_data ) @@ -291,6 +303,8 @@ async def parse(self, context: Context) -> dict: status_code=500, detail=f"Error during ingestion: {str(e)}", ) from e + finally: + reset_r2r_trace_context(trace_token) @orchestration_provider.failure() async def on_failure(self, context: Context) -> None: @@ -390,14 +404,18 @@ async def embed(self, context: Context) -> dict: document_info = DocumentResponse(**document_info_dict) extractions = context.step_output("ingest")["extractions"] - - embedding_generator = self.ingestion_service.embed_document( - extractions - ) - embeddings = [ - embedding.model_dump() - async for embedding in embedding_generator - ] + input_data = context.workflow_input()["request"] + with r2r_ingestion_trace_context( + input_data, + task_id=str(context.workflow_run_id()), + ): + embedding_generator = self.ingestion_service.embed_document( + extractions + ) + embeddings = [ + embedding.model_dump() + async for embedding in embedding_generator + ] await self.ingestion_service.update_document_status( document_info, status=IngestionStatus.STORING diff --git a/py/core/main/orchestration/simple/ingestion_workflow.py b/py/core/main/orchestration/simple/ingestion_workflow.py index d30e24b425..b19d4a204b 100644 --- a/py/core/main/orchestration/simple/ingestion_workflow.py +++ b/py/core/main/orchestration/simple/ingestion_workflow.py @@ -16,6 +16,11 @@ num_tokens, update_settings_from_dict, ) +from core.utils.observability import ( + build_r2r_ingestion_trace_context, + reset_r2r_trace_context, + set_r2r_trace_context, +) from ...services import IngestionService @@ -24,6 +29,9 @@ def simple_ingestion_factory(service: IngestionService): async def ingest_files(input_data): + trace_token = set_r2r_trace_context( + build_r2r_ingestion_trace_context(input_data) + ) document_info = None try: from core.base import IngestionStatus @@ -209,6 +217,8 @@ async def ingest_files(input_data): raise HTTPException( status_code=500, detail=f"Error during ingestion: {str(e)}" ) from e + finally: + reset_r2r_trace_context(trace_token) async def _ensure_collections_exists( service: IngestionService, @@ -284,6 +294,9 @@ async def _ensure_collections_exists( raise e async def ingest_chunks(input_data): + trace_token = set_r2r_trace_context( + build_r2r_ingestion_trace_context(input_data) + ) document_info = None try: from core.base import IngestionStatus @@ -431,6 +444,8 @@ async def ingest_chunks(input_data): status_code=500, detail=f"Error during chunk ingestion: {str(e)}", ) from e + finally: + reset_r2r_trace_context(trace_token) async def update_chunk(input_data): from core.main import IngestionServiceAdapter diff --git a/py/core/main/services/ingestion_service.py b/py/core/main/services/ingestion_service.py index fccc7a33c7..82ad3ba7dd 100644 --- a/py/core/main/services/ingestion_service.py +++ b/py/core/main/services/ingestion_service.py @@ -30,6 +30,7 @@ VectorTableName, ) from core.base.api.models import User +from core.utils.observability import r2r_observation_context from shared.abstractions import PDFParsingError, PopplerNotFoundError from ..abstractions import R2RProviders @@ -316,22 +317,24 @@ async def augment_document_info( }, ) - response = await self.providers.llm.aget_completion( - messages=messages, - generation_config=GenerationConfig( - model=self.config.ingestion.document_summary_model - or self.config.app.fast_llm - ), - ) + with r2r_observation_context("R2R: Document summary"): + response = await self.providers.llm.aget_completion( + messages=messages, + generation_config=GenerationConfig( + model=self.config.ingestion.document_summary_model + or self.config.app.fast_llm + ), + ) document_info.summary = response.choices[0].message.content # type: ignore if not document_info.summary: raise ValueError("Expected a generated response.") - embedding = await self.providers.embedding.async_get_embedding( - text=document_info.summary, - ) + with r2r_observation_context("R2R: Document summary embedding"): + embedding = await self.providers.embedding.async_get_embedding( + text=document_info.summary, + ) document_info.summary_embedding = embedding return @@ -366,9 +369,16 @@ async def process_batch( for ex in batch ] # Retrieve embeddings in bulk - vectors = await self.providers.embedding.async_get_embeddings( - texts, # list of strings - ) + with r2r_observation_context( + "R2R: Chunk embedding", + metadata={ + "r2r_chunk_count": len(batch), + "r2r_chunk_ids": [str(chunk.id) for chunk in batch], + }, + ): + vectors = await self.providers.embedding.async_get_embeddings( + texts, # list of strings + ) # Zip them back together results = [] for raw_vector, extraction in zip(vectors, batch, strict=False): @@ -733,35 +743,45 @@ async def _get_enriched_chunk_text( ] try: # Obtain the updated text from the LLM - updated_chunk_text = ( - ( - await self.providers.llm.aget_completion( - messages=await self.providers.database.prompts_handler.get_message_payload( - task_prompt_name=chunk_enrichment_settings.chunk_enrichment_prompt, - task_inputs={ - "document_summary": document_summary or "None", - "chunk": chunk["text"], - "preceding_chunks": ( - "\n".join(preceding_chunks) - if preceding_chunks - else "None" - ), - "succeeding_chunks": ( - "\n".join(succeeding_chunks) - if succeeding_chunks - else "None" - ), - "chunk_size": self.config.ingestion.chunk_size - or 1024, - }, - ), - generation_config=chunk_enrichment_settings.generation_config - or GenerationConfig(model=self.config.app.fast_llm), + with r2r_observation_context( + "R2R: Chunk enrichment", + metadata={ + "r2r_chunk_id": str(chunk["id"]), + "r2r_chunk_index": chunk_idx, + }, + ): + updated_chunk_text = ( + ( + await self.providers.llm.aget_completion( + messages=await self.providers.database.prompts_handler.get_message_payload( + task_prompt_name=chunk_enrichment_settings.chunk_enrichment_prompt, + task_inputs={ + "document_summary": document_summary + or "None", + "chunk": chunk["text"], + "preceding_chunks": ( + "\n".join(preceding_chunks) + if preceding_chunks + else "None" + ), + "succeeding_chunks": ( + "\n".join(succeeding_chunks) + if succeeding_chunks + else "None" + ), + "chunk_size": self.config.ingestion.chunk_size + or 1024, + }, + ), + generation_config=chunk_enrichment_settings.generation_config + or GenerationConfig( + model=self.config.app.fast_llm + ), + ) ) + .choices[0] + .message.content ) - .choices[0] - .message.content - ) except Exception: updated_chunk_text = chunk["text"] chunk["metadata"]["chunk_enrichment_status"] = "failed" @@ -775,9 +795,16 @@ async def _get_enriched_chunk_text( chunk["metadata"]["chunk_enrichment_status"] = "failed" # Re-embed - data = await self.providers.embedding.async_get_embedding( - updated_chunk_text - ) + with r2r_observation_context( + "R2R: Enriched chunk embedding", + metadata={ + "r2r_chunk_id": str(chunk["id"]), + "r2r_chunk_index": chunk_idx, + }, + ): + data = await self.providers.embedding.async_get_embedding( + updated_chunk_text + ) chunk["metadata"]["original_text"] = chunk["text"] return VectorEntry( diff --git a/py/core/main/services/retrieval_service.py b/py/core/main/services/retrieval_service.py index 29332cfdec..4843288fe1 100644 --- a/py/core/main/services/retrieval_service.py +++ b/py/core/main/services/retrieval_service.py @@ -48,6 +48,10 @@ find_new_citation_spans, num_tokens_from_messages, ) +from core.utils.observability import ( + r2r_observation_context, + r2r_trace_context, +) from shared.api.models.management.responses import MessageResponse from ..abstractions import R2RProviders @@ -269,14 +273,18 @@ async def search( an AggregateSearchResult that includes chunk + graph results. """ strategy = search_settings.search_strategy.lower() - - if strategy == "hyde": - return await self._hyde_search(query, search_settings) - elif strategy == "rag_fusion": - return await self._rag_fusion_search(query, search_settings) - else: - # 'vanilla', 'basic', or anything else... - return await self._basic_search(query, search_settings) + with r2r_trace_context( + "R2R: Search", + tags=("r2r-search",), + metadata={"search_strategy": strategy}, + ): + if strategy == "hyde": + return await self._hyde_search(query, search_settings) + elif strategy == "rag_fusion": + return await self._rag_fusion_search(query, search_settings) + else: + # 'vanilla', 'basic', or anything else... + return await self._basic_search(query, search_settings) async def _basic_search( self, query: str, search_settings: SearchSettings @@ -293,11 +301,10 @@ async def _basic_search( search_settings.use_semantic_search or search_settings.use_hybrid_search ): - query_vector = ( - await self.providers.completion_embedding.async_get_embedding( + with r2r_observation_context("R2R: Query embedding"): + query_vector = await self.providers.completion_embedding.async_get_embedding( text=query ) - ) # -- 2) Chunk search chunk_results = [] @@ -433,10 +440,11 @@ async def _generate_similar_queries( temperature=0.8, stream=False, ) - response = await self.providers.llm.aget_completion( - messages=[{"role": "system", "content": prompt}], - generation_config=gen_config, - ) + with r2r_observation_context("R2R: Search query generation"): + response = await self.providers.llm.aget_completion( + messages=[{"role": "system", "content": prompt}], + generation_config=gen_config, + ) raw_text = ( response.choices[0].message.content.strip() if response.choices[0].message.content is not None @@ -624,9 +632,12 @@ async def _fanout_chunk_and_graph_search( 2) chunk search + graph search with that embedding """ # Precompute the embedding of alt_text - vec = await self.providers.completion_embedding.async_get_embedding( - text=alt_text - ) + with r2r_observation_context("R2R: HyDE document embedding"): + vec = ( + await self.providers.completion_embedding.async_get_embedding( + text=alt_text + ) + ) # chunk search chunk_results = [] @@ -669,11 +680,10 @@ async def _vector_search_logic( search_settings.use_semantic_search or search_settings.use_hybrid_search ): - query_vector = ( - await self.providers.completion_embedding.async_get_embedding( + with r2r_observation_context("R2R: Query embedding"): + query_vector = await self.providers.completion_embedding.async_get_embedding( text=query_text ) - ) # 2) Choose which search to run if ( @@ -749,11 +759,10 @@ async def _graph_search_logic( # 1) Possibly embed query_embedding = precomputed_vector if query_embedding is None: - query_embedding = ( - await self.providers.completion_embedding.async_get_embedding( + with r2r_observation_context("R2R: Graph query embedding"): + query_embedding = await self.providers.completion_embedding.async_get_embedding( query_text ) - ) base_limit = search_settings.limit graph_limits = search_settings.graph_settings.limits or {} @@ -922,10 +931,11 @@ async def _run_hyde_generation( stream=False, ) - response = await self.providers.llm.aget_completion( - messages=[{"role": "system", "content": hyde_template}], - generation_config=completion_config, - ) + with r2r_observation_context("R2R: HyDE generation"): + response = await self.providers.llm.aget_completion( + messages=[{"role": "system", "content": hyde_template}], + generation_config=completion_config, + ) # Suppose the LLM returns something like: # @@ -946,11 +956,10 @@ async def search_documents( query_embedding: Optional[list[float]] = None, ) -> list[DocumentResponse]: if query_embedding is None: - query_embedding = ( - await self.providers.completion_embedding.async_get_embedding( + with r2r_observation_context("R2R: Document search embedding"): + query_embedding = await self.providers.completion_embedding.async_get_embedding( query ) - ) return ( await self.providers.database.documents_handler.search_documents( @@ -967,20 +976,24 @@ async def completion( *args, **kwargs, ): - return await self.providers.llm.aget_completion( - [message.to_dict() for message in messages], # type: ignore - generation_config, - *args, - **kwargs, - ) + with r2r_trace_context("R2R: Completion"): + with r2r_observation_context("R2R: Completion"): + return await self.providers.llm.aget_completion( + [message.to_dict() for message in messages], # type: ignore + generation_config, + *args, + **kwargs, + ) async def embedding( self, text: str, ): - return await self.providers.completion_embedding.async_get_embedding( - text=text - ) + with r2r_trace_context("R2R: Embedding"): + with r2r_observation_context("R2R: Embedding"): + return await self.providers.completion_embedding.async_get_embedding( + text=text + ) async def rag( self, @@ -1042,10 +1055,11 @@ async def rag( # 5) Check streaming vs. non-streaming if not rag_generation_config.stream: # ========== Non-Streaming Logic ========== - response = await self.providers.llm.aget_completion( - messages=messages, - generation_config=rag_generation_config, - ) + with r2r_observation_context("R2R: RAG response"): + response = await self.providers.llm.aget_completion( + messages=messages, + generation_config=rag_generation_config, + ) llm_text = response.choices[0].message.content # (a) Extract short-ID references from final text diff --git a/py/core/parsers/media/audio_parser.py b/py/core/parsers/media/audio_parser.py index 7d5f9f1d4d..827b6a487e 100644 --- a/py/core/parsers/media/audio_parser.py +++ b/py/core/parsers/media/audio_parser.py @@ -12,6 +12,7 @@ DatabaseProvider, IngestionConfig, ) +from core.utils.observability import build_litellm_metadata logger = logging.getLogger() @@ -52,12 +53,19 @@ async def ingest( # type: ignore temp_file_path = temp_file.name # Call Whisper transcription - response = await self.atranscription( - model=self.config.audio_transcription_model - or self.config.app.audio_lm, - file=open(temp_file_path, "rb"), - **kwargs, + metadata = build_litellm_metadata( + kwargs.pop("metadata", None), + default_trace_name="R2R: Audio transcription", + default_generation_name="R2R: Audio transcription", ) + with open(temp_file_path, "rb") as audio_file: + response = await self.atranscription( + model=self.config.audio_transcription_model + or self.config.app.audio_lm, + file=audio_file, + metadata=metadata, + **kwargs, + ) # The response should contain the transcribed text directly yield response.text diff --git a/py/core/parsers/media/img_parser.py b/py/core/parsers/media/img_parser.py index 1bf64ee735..bf321e6b9f 100644 --- a/py/core/parsers/media/img_parser.py +++ b/py/core/parsers/media/img_parser.py @@ -305,7 +305,9 @@ async def ingest( ] response = await self.llm_provider.aget_completion( - messages=messages, generation_config=generation_config + messages=messages, + generation_config=generation_config, + metadata={"generation_name": "R2R: Image extraction"}, ) if not response.choices or not response.choices[0].message: diff --git a/py/core/parsers/media/pdf_parser.py b/py/core/parsers/media/pdf_parser.py index dc4ca8eec1..88e5189d59 100644 --- a/py/core/parsers/media/pdf_parser.py +++ b/py/core/parsers/media/pdf_parser.py @@ -156,6 +156,10 @@ async def process_page(self, image, page_num: int) -> dict[str, str]: messages=messages, generation_config=generation_config, apply_timeout=True, + metadata={ + "generation_name": "R2R: PDF page extraction", + "r2r_page_number": page_num, + }, tools=[ { "name": "parse_pdf_page", @@ -202,6 +206,10 @@ async def process_page(self, image, page_num: int) -> dict[str, str]: messages=messages, generation_config=generation_config, apply_timeout=True, + metadata={ + "generation_name": "R2R: PDF page extraction", + "r2r_page_number": page_num, + }, ) if response.choices and response.choices[0].message: diff --git a/py/core/providers/embeddings/litellm.py b/py/core/providers/embeddings/litellm.py index 7322d7de4a..41c54dff09 100644 --- a/py/core/providers/embeddings/litellm.py +++ b/py/core/providers/embeddings/litellm.py @@ -16,6 +16,7 @@ EmbeddingProvider, R2RException, ) +from core.utils.observability import build_litellm_metadata from .utils import truncate_texts_to_token_limit @@ -74,6 +75,11 @@ def _get_embedding_kwargs(self, **kwargs): if self.config.api_key: embedding_kwargs["api_key"] = self.config.api_key embedding_kwargs.update(kwargs) + embedding_kwargs["metadata"] = build_litellm_metadata( + embedding_kwargs.get("metadata"), + default_trace_name="R2R: Embedding", + default_generation_name="R2R: Embedding", + ) return embedding_kwargs async def _execute_task(self, task: dict[str, Any]) -> list[list[float]]: diff --git a/py/core/providers/llm/litellm.py b/py/core/providers/llm/litellm.py index 44d467c2aa..9bd03c41c6 100644 --- a/py/core/providers/llm/litellm.py +++ b/py/core/providers/llm/litellm.py @@ -6,6 +6,7 @@ from core.base.abstractions import GenerationConfig from core.base.providers.llm import CompletionConfig, CompletionProvider +from core.utils.observability import build_litellm_metadata logger = logging.getLogger() @@ -53,6 +54,11 @@ async def _execute_task(self, task: dict[str, Any]): args = self._get_base_args(generation_config) args["messages"] = messages args = {**args, **kwargs} + args["metadata"] = build_litellm_metadata( + args.get("metadata"), + default_trace_name="R2R: LLM completion", + default_generation_name="R2R: LLM completion", + ) logger.debug( f"Executing LiteLLM task with generation_config={generation_config}" @@ -68,6 +74,11 @@ def _execute_task_sync(self, task: dict[str, Any]): args = self._get_base_args(generation_config) args["messages"] = messages args = {**args, **kwargs} + args["metadata"] = build_litellm_metadata( + args.get("metadata"), + default_trace_name="R2R: LLM completion", + default_generation_name="R2R: LLM completion", + ) logger.debug( f"Executing LiteLLM task with generation_config={generation_config}" diff --git a/py/core/utils/observability.py b/py/core/utils/observability.py new file mode 100644 index 0000000000..3b6c4e8900 --- /dev/null +++ b/py/core/utils/observability.py @@ -0,0 +1,204 @@ +from contextlib import contextmanager +from contextvars import ContextVar, Token +from dataclasses import dataclass, field +from typing import Any, Iterator, Mapping +from uuid import uuid4 + +R2R_TAG = "r2r" + + +@dataclass(frozen=True) +class R2RTraceContext: + trace_id: str + trace_name: str + tags: tuple[str, ...] = (R2R_TAG,) + metadata: Mapping[str, Any] = field(default_factory=dict) + session_id: str | None = None + + +@dataclass(frozen=True) +class R2RObservationContext: + generation_name: str + metadata: Mapping[str, Any] = field(default_factory=dict) + + +_trace_context: ContextVar[R2RTraceContext | None] = ContextVar( + "r2r_trace_context", default=None +) +_observation_context: ContextVar[R2RObservationContext | None] = ContextVar( + "r2r_observation_context", default=None +) + + +def create_r2r_trace_id() -> str: + """Create a Langfuse-compatible 32-character hexadecimal trace ID.""" + return uuid4().hex + + +def get_r2r_trace_context() -> R2RTraceContext | None: + return _trace_context.get() + + +def get_r2r_observation_context() -> R2RObservationContext | None: + return _observation_context.get() + + +def set_r2r_trace_context(context: R2RTraceContext) -> Token: + return _trace_context.set(context) + + +def reset_r2r_trace_context(token: Token) -> None: + _trace_context.reset(token) + + +def _unique_tags(*tag_groups: object) -> list[str]: + tags: list[str] = [] + for group in tag_groups: + values = [group] if isinstance(group, str) else group + if not isinstance(values, (list, tuple, set, frozenset)): + continue + for value in values: + if isinstance(value, str) and value not in tags: + tags.append(value) + return tags + + +@contextmanager +def r2r_trace_context( + trace_name: str, + *, + trace_id: str | None = None, + tags: tuple[str, ...] = (), + metadata: Mapping[str, Any] | None = None, + session_id: str | None = None, +) -> Iterator[R2RTraceContext]: + """Set trace-level attributes inherited by nested R2R model calls.""" + current = get_r2r_trace_context() + context = R2RTraceContext( + trace_id=(current.trace_id if current else trace_id) + or create_r2r_trace_id(), + trace_name=current.trace_name if current else trace_name, + tags=tuple( + _unique_tags( + current.tags if current else (), + (R2R_TAG,), + tags, + ) + ), + metadata={ + **(dict(current.metadata) if current else {}), + **dict(metadata or {}), + }, + session_id=(current.session_id if current else None) or session_id, + ) + token = _trace_context.set(context) + try: + yield context + finally: + _trace_context.reset(token) + + +@contextmanager +def r2r_observation_context( + generation_name: str, + *, + metadata: Mapping[str, Any] | None = None, +) -> Iterator[R2RObservationContext]: + """Set operation-level attributes for one nested R2R model call.""" + context = R2RObservationContext( + generation_name=generation_name, + metadata=dict(metadata or {}), + ) + token = _observation_context.set(context) + try: + yield context + finally: + _observation_context.reset(token) + + +@contextmanager +def r2r_ingestion_trace_context( + input_data: Mapping[str, Any], + *, + task_id: str | None = None, +) -> Iterator[R2RTraceContext]: + """Create the shared trace used by one document-ingestion attempt.""" + context = build_r2r_ingestion_trace_context( + input_data, + task_id=task_id, + ) + token = set_r2r_trace_context(context) + try: + yield context + finally: + reset_r2r_trace_context(token) + + +def build_r2r_ingestion_trace_context( + input_data: Mapping[str, Any], + *, + task_id: str | None = None, +) -> R2RTraceContext: + """Build the trace attributes for one document-ingestion attempt.""" + document_id = str(input_data["document_id"]) + trace_metadata = {"document_id": document_id} + if task_id: + trace_metadata["r2r_task_id"] = task_id + supplied_trace_id = input_data.get("langfuse_trace_id") + return R2RTraceContext( + trace_id=( + str(supplied_trace_id) + if supplied_trace_id + else create_r2r_trace_id() + ), + trace_name="R2R: Document ingestion", + tags=(R2R_TAG, "r2r-ingestion"), + metadata=trace_metadata, + ) + + +def build_litellm_metadata( + existing: Mapping[str, Any] | None = None, + *, + default_trace_name: str, + default_generation_name: str, +) -> dict[str, Any]: + """Merge R2R trace context into LiteLLM's Langfuse metadata fields.""" + metadata = dict(existing or {}) + trace = get_r2r_trace_context() + observation = get_r2r_observation_context() + + metadata["tags"] = _unique_tags( + metadata.get("tags", ()), + trace.tags if trace else (), + (R2R_TAG,), + ) + metadata.setdefault( + "trace_name", trace.trace_name if trace else default_trace_name + ) + metadata.setdefault( + "generation_name", + ( + observation.generation_name + if observation + else default_generation_name + ), + ) + + if trace: + metadata.setdefault("trace_id", trace.trace_id) + if trace.session_id: + metadata.setdefault("session_id", trace.session_id) + + trace_metadata = dict(trace.metadata) + existing_trace_metadata = metadata.get("trace_metadata") + if isinstance(existing_trace_metadata, Mapping): + trace_metadata.update(existing_trace_metadata) + if trace_metadata: + metadata["trace_metadata"] = trace_metadata + + if observation: + for key, value in observation.metadata.items(): + metadata.setdefault(key, value) + + return metadata diff --git a/py/tests/unit/test_litellm_observability.py b/py/tests/unit/test_litellm_observability.py new file mode 100644 index 0000000000..c8c7d275ff --- /dev/null +++ b/py/tests/unit/test_litellm_observability.py @@ -0,0 +1,62 @@ +from unittest.mock import AsyncMock + +import pytest +from core.base import EmbeddingConfig, GenerationConfig +from core.base.providers.llm import CompletionConfig +from core.providers.embeddings.litellm import LiteLLMEmbeddingProvider +from core.providers.llm.litellm import LiteLLMCompletionProvider +from core.utils.observability import ( + r2r_observation_context, + r2r_trace_context, +) + + +@pytest.mark.asyncio +async def test_completion_provider_adds_current_langfuse_context(): + provider = LiteLLMCompletionProvider(CompletionConfig(provider="litellm")) + provider.acompletion = AsyncMock(return_value="completion") + + with r2r_trace_context( + "R2R: Document ingestion", + trace_id="a" * 32, + tags=("r2r-ingestion",), + metadata={"document_id": "document-123"}, + ): + with r2r_observation_context("R2R: Document summary"): + result = await provider._execute_task( + { + "messages": [{"role": "user", "content": "Summarize"}], + "generation_config": GenerationConfig(model="test-model"), + "kwargs": {}, + } + ) + + assert result == "completion" + metadata = provider.acompletion.await_args.kwargs["metadata"] + assert metadata == { + "tags": ["r2r", "r2r-ingestion"], + "trace_name": "R2R: Document ingestion", + "generation_name": "R2R: Document summary", + "trace_id": "a" * 32, + "trace_metadata": {"document_id": "document-123"}, + } + + +def test_embedding_provider_adds_r2r_fallback_metadata(): + provider = LiteLLMEmbeddingProvider( + EmbeddingConfig( + provider="litellm", + base_model="test-embedding-model", + base_dimension=3, + ) + ) + + kwargs = provider._get_embedding_kwargs( + metadata={"tags": ["existing-tag"]} + ) + + assert kwargs["metadata"] == { + "tags": ["existing-tag", "r2r"], + "trace_name": "R2R: Embedding", + "generation_name": "R2R: Embedding", + } diff --git a/py/tests/unit/test_observability.py b/py/tests/unit/test_observability.py new file mode 100644 index 0000000000..4d159272d2 --- /dev/null +++ b/py/tests/unit/test_observability.py @@ -0,0 +1,115 @@ +from core.utils.observability import ( + build_litellm_metadata, + build_r2r_ingestion_trace_context, + get_r2r_trace_context, + r2r_observation_context, + r2r_trace_context, +) + + +def test_builds_ingestion_trace_from_workflow_input(): + trace = build_r2r_ingestion_trace_context( + { + "document_id": "document-123", + "langfuse_trace_id": "a" * 32, + }, + task_id="task-456", + ) + + assert trace.trace_id == "a" * 32 + assert trace.trace_name == "R2R: Document ingestion" + assert trace.tags == ("r2r", "r2r-ingestion") + assert trace.metadata == { + "document_id": "document-123", + "r2r_task_id": "task-456", + } + + +def test_merges_trace_and_observation_context_into_litellm_metadata(): + with r2r_trace_context( + "R2R: Document ingestion", + trace_id="b" * 32, + tags=("r2r-ingestion",), + metadata={"document_id": "document-123"}, + ): + with r2r_observation_context( + "R2R: Chunk enrichment", + metadata={"r2r_chunk_id": "chunk-789"}, + ): + metadata = build_litellm_metadata( + { + "tags": ["existing-tag"], + "trace_metadata": {"request_source": "api"}, + "custom": "value", + }, + default_trace_name="R2R: LLM completion", + default_generation_name="R2R: LLM completion", + ) + + assert metadata == { + "tags": ["existing-tag", "r2r", "r2r-ingestion"], + "trace_name": "R2R: Document ingestion", + "generation_name": "R2R: Chunk enrichment", + "trace_id": "b" * 32, + "trace_metadata": { + "document_id": "document-123", + "request_source": "api", + }, + "r2r_chunk_id": "chunk-789", + "custom": "value", + } + + +def test_preserves_explicit_litellm_names_and_trace_id(): + with r2r_trace_context( + "R2R: Search", + trace_id="c" * 32, + tags=("r2r-search",), + ): + metadata = build_litellm_metadata( + { + "trace_id": "d" * 32, + "trace_name": "Caller trace", + "generation_name": "Caller generation", + }, + default_trace_name="R2R: LLM completion", + default_generation_name="R2R: LLM completion", + ) + + assert metadata["trace_id"] == "d" * 32 + assert metadata["trace_name"] == "Caller trace" + assert metadata["generation_name"] == "Caller generation" + assert metadata["tags"] == ["r2r", "r2r-search"] + + +def test_nested_trace_context_keeps_workflow_trace_and_restores_it(): + assert get_r2r_trace_context() is None + + with r2r_trace_context( + "R2R: Document ingestion", + trace_id="e" * 32, + tags=("r2r-ingestion",), + metadata={"document_id": "document-123"}, + ): + with r2r_trace_context( + "R2R: Search", + tags=("r2r-search",), + metadata={"search_strategy": "basic"}, + ) as nested: + assert nested.trace_id == "e" * 32 + assert nested.trace_name == "R2R: Document ingestion" + assert nested.tags == ( + "r2r", + "r2r-ingestion", + "r2r-search", + ) + assert nested.metadata == { + "document_id": "document-123", + "search_strategy": "basic", + } + + assert get_r2r_trace_context().trace_name == ( + "R2R: Document ingestion" + ) + + assert get_r2r_trace_context() is None