Skip to content
Draft
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
297 changes: 94 additions & 203 deletions sentry_sdk/integrations/langchain.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@
import sentry_sdk
from sentry_sdk.ai.utils import (
GEN_AI_ALLOWED_MESSAGE_ROLES,
get_start_span_function,
normalize_message_roles,
set_data_normalized,
transform_content_part,
Expand Down Expand Up @@ -263,7 +262,7 @@ class SentryLangchainCallback(BaseCallbackHandler):
"""Callback handler that creates Sentry spans."""

def __init__(self, include_prompts: bool) -> None:
self.span_map: "OrderedDict[UUID, Union[sentry_sdk.tracing.Span, StreamedSpan]]" = OrderedDict()
self.span_map: "OrderedDict[UUID, StreamedSpan]" = OrderedDict()
self.include_prompts = include_prompts

def _handle_error(self, run_id: "UUID", error: "Any") -> None:
Expand Down Expand Up @@ -299,12 +298,10 @@ def _create_span(
op: str,
name: str,
origin: str,
) -> "Union[sentry_sdk.tracing.Span, StreamedSpan]":
) -> "StreamedSpan":
span = None
if parent_id:
parent_span: "Optional[Union[sentry_sdk.tracing.Span, StreamedSpan]]" = (
self.span_map.get(parent_id)
)
parent_span: "Optional[StreamedSpan]" = self.span_map.get(parent_id)
if parent_span:
span = (
sentry_sdk.traces.start_span(
Expand All @@ -320,17 +317,12 @@ def _create_span(
)

if span is None:
span_streaming = has_span_streaming_enabled(sentry_sdk.get_client().options)
span = (
sentry_sdk.traces.start_span(
name=name,
attributes={
"sentry.op": op,
"sentry.origin": origin,
},
)
if span_streaming
else sentry_sdk.start_span(op=op, name=name, origin=origin)
span = sentry_sdk.traces.start_span(
name=name,
attributes={
"sentry.op": op,
"sentry.origin": origin,
},
)

span.__enter__()
Expand All @@ -339,7 +331,7 @@ def _create_span(

def _exit_span(
self: "SentryLangchainCallback",
span: "Union[sentry_sdk.tracing.Span, StreamedSpan]",
span: "StreamedSpan",
run_id: "UUID",
) -> None:
span.__exit__(None, None, None)
Expand Down Expand Up @@ -1157,89 +1149,46 @@ def new_invoke(self: "Any", *args: "Any", **kwargs: "Any") -> "Any":
record_inputs = True
record_outputs = True

if has_span_streaming_enabled(client.options):
with sentry_sdk.traces.start_span(
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
attributes={
"sentry.op": OP.GEN_AI_INVOKE_AGENT,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "invoke_agent",
SPANDATA.GEN_AI_RESPONSE_STREAMING: False,
},
) as span:
if run_name:
span.set_attribute(SPANDATA.GEN_AI_FUNCTION_ID, run_name)

_set_tools_on_span(span, tools)

# Run the agent
result = f(self, *args, **kwargs)

input = result.get("input")
if input is not None and record_inputs:
normalized_messages = normalize_message_roles([input])

scope = sentry_sdk.get_current_scope()
messages_data = (
truncate_and_annotate_messages(normalized_messages, span, scope)
if not has_span_streaming_enabled(client.options)
else normalized_messages
)
if messages_data is not None:
set_data_normalized(
span,
SPANDATA.GEN_AI_REQUEST_MESSAGES,
messages_data,
unpack=False,
)

output = result.get("output")
if output is not None and record_outputs:
set_data_normalized(span, SPANDATA.GEN_AI_RESPONSE_TEXT, output)

return result
else:
start_span_function = get_start_span_function()

with start_span_function(
op=OP.GEN_AI_INVOKE_AGENT,
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
origin=LangchainIntegration.origin,
) as span:
if run_name:
span.set_data(SPANDATA.GEN_AI_FUNCTION_ID, run_name)

span.set_data(SPANDATA.GEN_AI_OPERATION_NAME, "invoke_agent")
span.set_data(SPANDATA.GEN_AI_RESPONSE_STREAMING, False)
with sentry_sdk.traces.start_span(
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
attributes={
"sentry.op": OP.GEN_AI_INVOKE_AGENT,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "invoke_agent",
SPANDATA.GEN_AI_RESPONSE_STREAMING: False,
},
) as span:
if run_name:
span.set_attribute(SPANDATA.GEN_AI_FUNCTION_ID, run_name)

_set_tools_on_span(span, tools)
_set_tools_on_span(span, tools)

# Run the agent
result = f(self, *args, **kwargs)
# Run the agent
result = f(self, *args, **kwargs)

input = result.get("input")
if input is not None and record_inputs:
normalized_messages = normalize_message_roles([input])
input = result.get("input")
if input is not None and record_inputs:
normalized_messages = normalize_message_roles([input])

scope = sentry_sdk.get_current_scope()
messages_data = (
truncate_and_annotate_messages(normalized_messages, span, scope)
if not has_span_streaming_enabled(client.options)
else normalized_messages
scope = sentry_sdk.get_current_scope()
messages_data = (
truncate_and_annotate_messages(normalized_messages, span, scope)
if not has_span_streaming_enabled(client.options)
else normalized_messages
)
if messages_data is not None:
set_data_normalized(
span,
SPANDATA.GEN_AI_REQUEST_MESSAGES,
messages_data,
unpack=False,
)
if messages_data is not None:
set_data_normalized(
span,
SPANDATA.GEN_AI_REQUEST_MESSAGES,
messages_data,
unpack=False,
)

output = result.get("output")
if output is not None and record_outputs:
set_data_normalized(span, SPANDATA.GEN_AI_RESPONSE_TEXT, output)
output = result.get("output")
if output is not None and record_outputs:
set_data_normalized(span, SPANDATA.GEN_AI_RESPONSE_TEXT, output)

return result
return result

return new_invoke

Expand All @@ -1264,34 +1213,18 @@ def new_stream(self: "Any", *args: "Any", **kwargs: "Any") -> "Any":
record_inputs = True
record_outputs = True

if has_span_streaming_enabled(client.options):
span = sentry_sdk.traces.start_span(
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
attributes={
"sentry.op": OP.GEN_AI_INVOKE_AGENT,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "invoke_agent",
SPANDATA.GEN_AI_RESPONSE_STREAMING: True,
},
)

if run_name:
span.set_attribute(SPANDATA.GEN_AI_FUNCTION_ID, run_name)
else:
start_span_function = get_start_span_function()

span = start_span_function(
op=OP.GEN_AI_INVOKE_AGENT,
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
origin=LangchainIntegration.origin,
)
span.__enter__()

span.set_data(SPANDATA.GEN_AI_OPERATION_NAME, "invoke_agent")
span.set_data(SPANDATA.GEN_AI_RESPONSE_STREAMING, True)
span = sentry_sdk.traces.start_span(
name=f"invoke_agent {run_name}" if run_name else "invoke_agent",
attributes={
"sentry.op": OP.GEN_AI_INVOKE_AGENT,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "invoke_agent",
SPANDATA.GEN_AI_RESPONSE_STREAMING: True,
},
)

if run_name:
span.set_data(SPANDATA.GEN_AI_FUNCTION_ID, run_name)
if run_name:
span.set_attribute(SPANDATA.GEN_AI_FUNCTION_ID, run_name)

_set_tools_on_span(span, tools)

Expand Down Expand Up @@ -1410,48 +1343,27 @@ def new_embedding_method(self: "Any", *args: "Any", **kwargs: "Any") -> "Any":
# TODO: Remove this branch once `send_default_pii` is deprecated
record_inputs = True

if has_span_streaming_enabled(client.options):
with sentry_sdk.traces.start_span(
name=f"embeddings {model_name}" if model_name else "embeddings",
attributes={
"sentry.op": OP.GEN_AI_EMBEDDINGS,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "embeddings",
},
) as span:
if model_name:
span.set_attribute(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)

result = f(self, *args, **kwargs)
return result
else:
with sentry_sdk.start_span(
op=OP.GEN_AI_EMBEDDINGS,
name=f"embeddings {model_name}" if model_name else "embeddings",
origin=LangchainIntegration.origin,
) as span:
span.set_data(SPANDATA.GEN_AI_OPERATION_NAME, "embeddings")
if model_name:
span.set_data(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)
with sentry_sdk.traces.start_span(
name=f"embeddings {model_name}" if model_name else "embeddings",
attributes={
"sentry.op": OP.GEN_AI_EMBEDDINGS,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "embeddings",
},
) as span:
if model_name:
span.set_attribute(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)

result = f(self, *args, **kwargs)
return result
result = f(self, *args, **kwargs)
return result

return new_embedding_method

Expand All @@ -1477,47 +1389,26 @@ async def new_async_embedding_method(
# TODO: Remove this branch once `send_default_pii` is deprecated
record_inputs = True

if has_span_streaming_enabled(client.options):
with sentry_sdk.traces.start_span(
name=f"embeddings {model_name}" if model_name else "embeddings",
attributes={
"sentry.op": OP.GEN_AI_EMBEDDINGS,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "embeddings",
},
) as span:
if model_name:
span.set_attribute(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)

result = await f(self, *args, **kwargs)
return result
else:
with sentry_sdk.start_span(
op=OP.GEN_AI_EMBEDDINGS,
name=f"embeddings {model_name}" if model_name else "embeddings",
origin=LangchainIntegration.origin,
) as span:
span.set_data(SPANDATA.GEN_AI_OPERATION_NAME, "embeddings")
if model_name:
span.set_data(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)
with sentry_sdk.traces.start_span(
name=f"embeddings {model_name}" if model_name else "embeddings",
attributes={
"sentry.op": OP.GEN_AI_EMBEDDINGS,
"sentry.origin": LangchainIntegration.origin,
SPANDATA.GEN_AI_OPERATION_NAME: "embeddings",
},
) as span:
if model_name:
span.set_attribute(SPANDATA.GEN_AI_REQUEST_MODEL, model_name)

if record_inputs and len(args) > 0:
input_data = args[0]
# Normalize to list format
texts = input_data if isinstance(input_data, list) else [input_data]
set_data_normalized(
span, SPANDATA.GEN_AI_EMBEDDINGS_INPUT, texts, unpack=False
)

result = await f(self, *args, **kwargs)
return result
result = await f(self, *args, **kwargs)
return result

return new_async_embedding_method
Loading
Loading