Skip to content
Merged
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
61 changes: 40 additions & 21 deletions sentry_sdk/integrations/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -483,19 +483,11 @@ def _set_completions_api_input_data(
kwargs: "dict[str, Any]",
integration: "OpenAIIntegration",
) -> None:
messages: "Optional[Union[str, Iterable[ChatCompletionMessageParam]]]" = kwargs.get(
"messages"
)

set_on_span = (
span.set_attribute if isinstance(span, StreamedSpan) else span.set_data
)
tools = kwargs.get("tools")
if tools is not None and _is_given(tools):
set_on_span(
SPANDATA.GEN_AI_TOOL_DEFINITIONS,
json.dumps(_transform_tool_definitions_completions(tools)),
)

set_data_normalized(span, SPANDATA.GEN_AI_OPERATION_NAME, "chat")

model = kwargs.get("model")
if model is not None:
Expand Down Expand Up @@ -525,17 +517,48 @@ def _set_completions_api_input_data(
if reasoning_level is not None and _is_given(reasoning_level):
set_on_span(SPANDATA.GEN_AI_REQUEST_REASONING_LEVEL, reasoning_level)

if (
not should_send_default_pii()
or not integration.include_prompts
or messages is None
):
set_data_normalized(span, SPANDATA.GEN_AI_OPERATION_NAME, "chat")
client = sentry_sdk.get_client()
if has_data_collection_enabled(client.options):
if (
integration.include_prompts
and client.options["data_collection"]["gen_ai"]["inputs"]
):
tools = kwargs.get("tools")
if tools is not None and _is_given(tools):
set_on_span(
SPANDATA.GEN_AI_TOOL_DEFINITIONS,
json.dumps(_transform_tool_definitions_completions(tools)),
)
else:
# Pre-data collection this was always set, so this needs to be left here for now until
# we deprecate `send_default_pii`. Once we do, this 'else' branch should be removed,
# and the above branch placed below the "if not should_send_default_pii() or not integration.include_prompts"
# line below
tools = kwargs.get("tools")
if tools is not None and _is_given(tools):
set_on_span(
SPANDATA.GEN_AI_TOOL_DEFINITIONS,
json.dumps(_transform_tool_definitions_completions(tools)),
)

messages: "Optional[Union[str, Iterable[ChatCompletionMessageParam]]]" = kwargs.get(
"messages"
)

if has_data_collection_enabled(client.options):
# This takes precedence over the global data collection settings
if not integration.include_prompts:
return
if not client.options["data_collection"]["gen_ai"]["inputs"]:
return
elif not should_send_default_pii() or not integration.include_prompts:
return

if messages is None:
return

if isinstance(messages, str):
normalized_messages = normalize_message_roles([messages]) # type: ignore
client = sentry_sdk.get_client()
scope = sentry_sdk.get_current_scope()
messages_data = (
truncate_and_annotate_messages(normalized_messages, span, scope)
Expand All @@ -546,12 +569,10 @@ def _set_completions_api_input_data(
set_data_normalized(
span, SPANDATA.GEN_AI_REQUEST_MESSAGES, messages_data, unpack=False
)
set_data_normalized(span, SPANDATA.GEN_AI_OPERATION_NAME, "chat")
return

# dict special case following https://github.com/openai/openai-python/blob/3e0c05b84a2056870abf3bd6a5e7849020209cc3/src/openai/_utils/_transform.py#L194-L197
if not isinstance(messages, Iterable) or isinstance(messages, dict):
set_data_normalized(span, SPANDATA.GEN_AI_OPERATION_NAME, "chat")
return

messages = list(messages)
Expand Down Expand Up @@ -583,8 +604,6 @@ def _set_completions_api_input_data(
span, SPANDATA.GEN_AI_REQUEST_MESSAGES, messages_data, unpack=False
)

set_data_normalized(span, SPANDATA.GEN_AI_OPERATION_NAME, "chat")


def _set_embeddings_input_data(
span: "Union[Span, StreamedSpan]",
Expand Down
168 changes: 168 additions & 0 deletions tests/integrations/openai/test_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,26 @@ async def __call__(self, *args, **kwargs):
}
]

EXAMPLE_COMPLETIONS_TOOLS = [
{
"type": "function",
"function": {
"name": "get_current_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
},
"required": ["location"],
},
},
}
]


@pytest.mark.skipif(
OPENAI_VERSION <= (2, 10, 0),
Expand Down Expand Up @@ -639,6 +659,154 @@ def test_nonstreaming_chat_completion(
assert span["data"]["gen_ai.usage.total_tokens"] == 30


@pytest.mark.skipif(
OPENAI_VERSION <= (1, 1, 0),
reason="OpenAI versions <=1.1.0 do not support the tools parameter.",
)
@pytest.mark.parametrize("span_streaming", [True, False])
@pytest.mark.parametrize("stream_gen_ai_spans", [True, False])
@pytest.mark.parametrize(
"data_collection,include_prompts,expected_present,expected_absent",
[
pytest.param(
{"gen_ai": {"inputs": True}},
True,
{
SPANDATA.GEN_AI_SYSTEM_INSTRUCTIONS: json.dumps(
[{"type": "text", "content": "You are a helpful assistant."}]
),
SPANDATA.GEN_AI_REQUEST_MESSAGES: safe_serialize(
[{"role": "user", "content": "hello"}]
),
SPANDATA.GEN_AI_TOOL_DEFINITIONS: safe_serialize(EXAMPLE_TOOLS),
},
[],
id="inputs-enabled",
),
pytest.param(
{"gen_ai": {"inputs": False}},
True,
{},
[
SPANDATA.GEN_AI_SYSTEM_INSTRUCTIONS,
SPANDATA.GEN_AI_REQUEST_MESSAGES,
SPANDATA.GEN_AI_TOOL_DEFINITIONS,
],
id="inputs-disabled",
),
pytest.param(
{},
True,
{
SPANDATA.GEN_AI_SYSTEM_INSTRUCTIONS: json.dumps(
[{"type": "text", "content": "You are a helpful assistant."}]
),
SPANDATA.GEN_AI_REQUEST_MESSAGES: safe_serialize(
[{"role": "user", "content": "hello"}]
),
SPANDATA.GEN_AI_TOOL_DEFINITIONS: safe_serialize(EXAMPLE_TOOLS),
},
[],
id="gen-ai-omitted-defaults-to-enabled",
),
pytest.param(
{"gen_ai": {"inputs": True}},
False,
{},
[
SPANDATA.GEN_AI_SYSTEM_INSTRUCTIONS,
SPANDATA.GEN_AI_REQUEST_MESSAGES,
SPANDATA.GEN_AI_TOOL_DEFINITIONS,
],
id="include-prompts-disabled-overrides-inputs-enabled",
),
],
)
def test_completions_api_data_collection(
sentry_init,
capture_events,
capture_items,
data_collection,
include_prompts,
expected_present,
expected_absent,
nonstreaming_chat_completions_model_response,
stream_gen_ai_spans,
span_streaming,
):
sentry_init(
integrations=[OpenAIIntegration(include_prompts=include_prompts)],
disabled_integrations=[StdlibIntegration],
traces_sample_rate=1.0,
_experiments={"data_collection": data_collection},
stream_gen_ai_spans=stream_gen_ai_spans,
trace_lifecycle="stream" if span_streaming else "static",
)

client = OpenAI(api_key="z")
client.chat.completions._post = mock.Mock(
return_value=nonstreaming_chat_completions_model_response(
response_id="chat-id",
response_model="gpt-3.5-turbo",
message_content="the model response",
created=10000000,
usage=CompletionUsage(
prompt_tokens=20,
completion_tokens=10,
total_tokens=30,
),
)
)

create_kwargs = {
"model": "some-model",
"messages": [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "hello"},
],
"max_tokens": 100,
"presence_penalty": 0.1,
"frequency_penalty": 0.2,
"temperature": 0.7,
"top_p": 0.9,
"tools": EXAMPLE_COMPLETIONS_TOOLS,
}

if span_streaming or stream_gen_ai_spans:
items = capture_items("span")

with start_transaction(name="openai tx"):
client.chat.completions.create(**create_kwargs)

sentry_sdk.flush()
(span,) = (item.payload for item in items)
span_data = span["attributes"]
else:
events = capture_events()

with start_transaction(name="openai tx"):
client.chat.completions.create(**create_kwargs)

(transaction,) = events
(span,) = transaction["spans"]
span_data = span["data"]

assert span_data[SPANDATA.GEN_AI_OPERATION_NAME] == "chat"
assert span_data[SPANDATA.GEN_AI_REQUEST_MODEL] == "some-model"
assert span_data[SPANDATA.GEN_AI_REQUEST_MAX_TOKENS] == 100
assert span_data[SPANDATA.GEN_AI_REQUEST_PRESENCE_PENALTY] == 0.1
assert span_data[SPANDATA.GEN_AI_REQUEST_FREQUENCY_PENALTY] == 0.2
assert span_data[SPANDATA.GEN_AI_REQUEST_TEMPERATURE] == 0.7
assert span_data[SPANDATA.GEN_AI_REQUEST_TOP_P] == 0.9
assert span_data[SPANDATA.GEN_AI_SYSTEM] == "openai"

for key, value in expected_present.items():
assert span_data[key] == value

for key in expected_absent:
assert key not in span_data


@pytest.mark.parametrize("span_streaming", [True, False])
@pytest.mark.parametrize("stream_gen_ai_spans", [True, False])
@pytest.mark.asyncio
Expand Down
Loading