Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Nest LangChain child spans (agent, workflow, chat, tool) under their caller instead of emitting each as its own root trace when callbacks run outside the ambient context (e.g. async LangGraph task loops).
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
LLMResult,
)

from opentelemetry.context import Context
from opentelemetry.instrumentation.genai.langchain.invocation_manager import (
_InvocationManager,
)
Expand All @@ -38,6 +39,7 @@
response_fields_from_generation,
to_input_messages,
)
from opentelemetry.trace import set_span_in_context
from opentelemetry.util.genai.handler import TelemetryHandler
from opentelemetry.util.genai.invocation import (
AgentInvocation,
Expand Down Expand Up @@ -86,6 +88,30 @@ def __init__(self, telemetry_handler: TelemetryHandler) -> None:
self._telemetry_handler = telemetry_handler
self._invocation_manager = _InvocationManager()

def _parent_context(self, parent_run_id: UUID | None) -> Context | None:
"""Return a context carrying the nearest emitted ancestor invocation's span.

LangChain runs queued callbacks in its own task loop where the
ambient OpenTelemetry context does not hold the parent span, and a
child's ``parent_run_id`` may point at an intermediate chain node that
emits no telemetry. Walk up the run hierarchy to the nearest live
invocation and parent to its span so child spans nest under their
caller instead of becoming their own root trace.

A parent is usable as long as its span context is valid even when the
span itself is not recorded (e.g. downstream sampling); parenting to it
keeps the trace tree's shape and depth correct.
"""
current = parent_run_id
visited: set[UUID] = set()
while current is not None and current not in visited:
visited.add(current)
parent = self._invocation_manager.get_invocation(current)
if parent is not None and parent.span.get_span_context().is_valid:
return set_span_in_context(parent.span)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The is_recording() gate discards valid non-recording parents and can produce orphan sampled children and incorrect span depth. Maybe use get_span_context().is_valid.

current = self._invocation_manager.get_parent_run_id(current)
return None

def on_chain_start(
self,
serialized: dict[str, Any],
Expand All @@ -108,7 +134,8 @@ def on_chain_start(
metadata.get("workflow_name") if metadata else None
)
workflow = self._telemetry_handler.workflow(
name=workflow_name_override or workflow_name
name=workflow_name_override or workflow_name,
parent_context=self._parent_context(parent_run_id),
)
workflow.conversation_id = conversation_id
if capture_content:
Expand Down Expand Up @@ -136,6 +163,7 @@ def on_chain_start(
if suggested_agent_name_lower != agent_invocation_name_lower:
agent = self._telemetry_handler.invoke_local_agent(
agent_name=suggested_agent_name,
parent_context=self._parent_context(parent_run_id),
)
agent.conversation_id = conversation_id
if capture_content:
Expand Down Expand Up @@ -300,6 +328,7 @@ def on_chat_model_start(
llm_invocation = self._telemetry_handler.inference(
provider,
request_model=request_model,
parent_context=self._parent_context(parent_run_id),
)
llm_invocation.conversation_id = _conversation_id(metadata)
llm_invocation.input_messages = input_messages
Expand Down Expand Up @@ -589,6 +618,7 @@ def on_tool_start(
tool_description=description,
tool_type="function",
agent_name=agent_name,
parent_context=self._parent_context(parent_run_id),
)
tool_invocation.arguments = arguments
tool_call_id = kwargs.get("tool_call_id")
Expand Down Expand Up @@ -647,7 +677,9 @@ def on_retriever_start(
provider = meta.get("ls_vector_store_provider") or None
request_model = meta.get("ls_embedding_model") or None
retrieval = self._telemetry_handler.retrieval(
provider=provider, request_model=request_model
provider=provider,
request_model=request_model,
parent_context=self._parent_context(parent_run_id),
)
retrieval.query_text = query
self._invocation_manager.add_invocation_state(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
LLMResult,
)

from opentelemetry import context
from opentelemetry.instrumentation.genai.langchain.callback_handler import (
OpenTelemetryLangChainCallbackHandler,
)
Expand All @@ -50,6 +51,7 @@
to_input_messages,
to_output_messages,
)
from opentelemetry.util.genai.handler import TelemetryHandler
from opentelemetry.util.genai.invocation import (
AgentInvocation,
InferenceInvocation,
Expand Down Expand Up @@ -169,7 +171,9 @@ def test_workflow_name_from_serialized(self):
parent_run_id=None,
)

telemetry.workflow.assert_called_once_with(name="MyLangGraph")
telemetry.workflow.assert_called_once_with(
name="MyLangGraph", parent_context=None
)

def test_workflow_name_overridden_by_metadata(self):
handler, telemetry, _, _ = _make_handler()
Expand All @@ -183,7 +187,9 @@ def test_workflow_name_overridden_by_metadata(self):
metadata={"workflow_name": "custom_workflow"},
)

telemetry.workflow.assert_called_once_with(name="custom_workflow")
telemetry.workflow.assert_called_once_with(
name="custom_workflow", parent_context=None
)

def test_workflow_conversation_id_from_metadata(self):
handler, _, workflow_inv, _ = _make_handler()
Expand Down Expand Up @@ -235,6 +241,7 @@ def test_new_agent_span_created(self):

telemetry.invoke_local_agent.assert_called_once_with(
agent_name="math_agent",
parent_context=None,
)
assert agent_inv.agent_name == "math_agent"
assert handler._invocation_manager.get_invocation(run_id) is agent_inv
Expand Down Expand Up @@ -1377,7 +1384,7 @@ def test_provider_passed_from_metadata(self):
)

telemetry.retrieval.assert_called_once_with(
provider="Chroma", request_model=None
provider="Chroma", request_model=None, parent_context=None
)

def test_provider_none_when_metadata_absent(self):
Expand All @@ -1391,7 +1398,7 @@ def test_provider_none_when_metadata_absent(self):
)

telemetry.retrieval.assert_called_once_with(
provider=None, request_model=None
provider=None, request_model=None, parent_context=None
)

def test_request_model_passed_from_ls_embedding_model(self):
Expand All @@ -1409,7 +1416,9 @@ def test_request_model_passed_from_ls_embedding_model(self):
)

telemetry.retrieval.assert_called_once_with(
provider="Chroma", request_model="text-embedding-3-small"
provider="Chroma",
request_model="text-embedding-3-small",
parent_context=None,
)

def test_request_model_none_when_ls_embedding_model_absent(self):
Expand All @@ -1424,7 +1433,7 @@ def test_request_model_none_when_ls_embedding_model_absent(self):
)

telemetry.retrieval.assert_called_once_with(
provider="Chroma", request_model=None
provider="Chroma", request_model=None, parent_context=None
)

def test_registered_in_invocation_manager(self):
Expand Down Expand Up @@ -2674,3 +2683,116 @@ def test_on_chat_model_start_captures_input_messages_when_content_enabled():
assert any(isinstance(p, TextPart) for p in parts)
blob = next(p for p in parts if isinstance(p, BlobPart))
assert blob.content == _REAL_PNG_BYTES


# ---------------------------------------------------------------------------
# Span parenting across callback types
# ---------------------------------------------------------------------------


class TestSpanParenting:
def test_chat_span_nested_under_agent_span(
self, tracer_provider, span_exporter
):
handler = OpenTelemetryLangChainCallbackHandler(
TelemetryHandler(tracer_provider=tracer_provider)
)

agent_run_id = _run_id()
chat_run_id = _run_id()

handler.on_chain_start(
serialized={"name": "agent_1"},
inputs={},
run_id=agent_run_id,
parent_run_id=None,
metadata={"agent_name": "agent_1", "ls_provider": "openai"},
)

# Mask the ambient context so the chat span can only be parented via
# the explicit parent_context resolved from parent_run_id.
mask = context.attach(context.Context())
try:
handler.on_chat_model_start(
serialized={},
messages=[[HumanMessage(content="hello")]],
run_id=chat_run_id,
parent_run_id=agent_run_id,
metadata={"model_name": "gpt-4o"},
)
finally:
context.detach(mask)

handler.on_llm_end(
response=LLMResult(
generations=[
[
ChatGeneration(
message=AIMessage(content="hi"),
generation_info={"finish_reason": "stop"},
)
]
]
),
run_id=chat_run_id,
)
handler.on_chain_end(outputs={}, run_id=agent_run_id)

spans = span_exporter.get_finished_spans()
chat = next(s for s in spans if s.name == "chat gpt-4o")
agent = next(s for s in spans if s.name == "invoke_agent agent_1")
assert chat.parent is not None
assert chat.parent.span_id == agent.context.span_id

def test_tool_span_nested_under_workflow_through_unemitted_node(
self, tracer_provider, span_exporter
):
handler = OpenTelemetryLangChainCallbackHandler(
TelemetryHandler(tracer_provider=tracer_provider)
)

workflow_run_id = _run_id()
# An intermediate chain node that emits no telemetry of its own but
# sits between the workflow and the tool in the run hierarchy.
intermediate_run_id = _run_id()
tool_run_id = _run_id()

handler.on_chain_start(
serialized={"name": "LangGraph"},
inputs={},
run_id=workflow_run_id,
parent_run_id=None,
)

# Register an intermediate chain node that emits no telemetry of its
# own but sits between the workflow and the tool in the hierarchy.
handler.on_chain_start(
serialized={"name": "intermediate_node"},
inputs={},
run_id=intermediate_run_id,
parent_run_id=workflow_run_id,
)

mask = context.attach(context.Context())
try:
handler.on_tool_start(
serialized=None,
input_str="{}",
run_id=tool_run_id,
parent_run_id=intermediate_run_id,
inputs={"x": 1},
)
finally:
context.detach(mask)

handler.on_tool_end(output={"x": 43}, run_id=tool_run_id)
handler.on_chain_end(outputs={}, run_id=workflow_run_id)

spans = span_exporter.get_finished_spans()
tool = next(s for s in spans if s.name == "execute_tool unknown")
workflow = next(
s for s in spans if s.name == "invoke_workflow LangGraph"
)
assert tool.parent is not None
assert tool.context.trace_id == workflow.context.trace_id
assert tool.parent.span_id == workflow.context.span_id
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from __future__ import annotations

from opentelemetry._logs import Logger
from opentelemetry.context import Context
from opentelemetry.semconv._incubating.attributes import (
gen_ai_attributes as GenAI,
)
Expand Down Expand Up @@ -49,6 +50,7 @@ def __init__(
server_address: str | None = None,
server_port: int | None = None,
agent_name: str | None = None,
parent_context: Context | None = None,
) -> None:
"""Use handler.invoke_local_agent() or handler.invoke_remote_agent() instead of calling this directly."""
_operation_name = GenAI.GenAiOperationNameValues.INVOKE_AGENT.value
Expand All @@ -62,6 +64,7 @@ def __init__(
if agent_name
else _operation_name,
span_kind=span_kind,
parent_context=parent_context,
)
self._provider: str | None = provider
self._request_model: str | None = request_model
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from dataclasses import dataclass, field

from opentelemetry._logs import Logger, LogRecord
from opentelemetry.context import Context
from opentelemetry.semconv._incubating.attributes import (
gen_ai_attributes as GenAI,
)
Expand Down Expand Up @@ -50,6 +51,7 @@ def __init__(
server_port: int | None = None,
operation_name: str | None = None,
error_type_resolver: ErrorTypeResolver | None = None,
parent_context: Context | None = None,
) -> None:
operation_name = (
operation_name or GenAI.GenAiOperationNameValues.CHAT.value
Expand All @@ -66,6 +68,7 @@ def __init__(
else operation_name,
span_kind=SpanKind.CLIENT,
error_type_resolver=error_type_resolver,
parent_context=parent_context,
)
self._provider: str = provider
self._request_model: str | None = request_model
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,12 +69,14 @@ def __init__(
attributes: dict[str, AttributeValue] | None = None,
metric_attributes: dict[str, AttributeValue] | None = None,
error_type_resolver: ErrorTypeResolver | None = None,
parent_context: Context | None = None,
) -> None:
self._tracer = tracer
self._metrics_recorder = metrics_recorder
self._logger = logger
self._completion_hook = completion_hook
self._error_type_resolver = error_type_resolver
self._parent_context = parent_context
self._operation_name: str = operation_name
self.attributes: dict[str, AttributeValue] = (
{} if attributes is None else attributes
Expand Down Expand Up @@ -105,13 +107,27 @@ def _start(

Args:
attributes: Initial span attributes available for sampling decisions.

The span is parented to ``self._parent_context`` when one was supplied
at construction; otherwise it parents from the ambient context.
"""
self.span = self._tracer.start_span(
name=self._span_name,
kind=self._span_kind,
attributes=attributes,
)
self._span_context = set_span_in_context(self.span)
if self._parent_context is not None:
self.span = self._tracer.start_span(
name=self._span_name,
kind=self._span_kind,
attributes=attributes,
context=self._parent_context,
)
self._span_context = set_span_in_context(
self.span, self._parent_context
)
else:
self.span = self._tracer.start_span(
name=self._span_name,
kind=self._span_kind,
attributes=attributes,
)
self._span_context = set_span_in_context(self.span)
self._monotonic_start_s = timeit.default_timer()
self._context_token = attach(self._span_context)

Expand Down
Loading