Skip to content

Commit ddb5dc6

Browse files
Add suppress_invoke_agent_input option to OpenAI trace processor
Co-authored-by: sergioescalera <8428450+sergioescalera@users.noreply.github.com>
1 parent 1b3b90b commit ddb5dc6

3 files changed

Lines changed: 325 additions & 5 deletions

File tree

libraries/microsoft-agents-a365-observability-extensions-openai/microsoft_agents_a365/observability/extensions/openai/trace_instrumentor.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@ def _instrument(self, **kwargs: Any) -> None:
4747
"""Instruments the OpenAI Agents SDK with Microsoft Agent 365 tracing."""
4848
tracer_name = kwargs["tracer_name"] if kwargs.get("tracer_name") else None
4949
tracer_version = kwargs["tracer_version"] if kwargs.get("tracer_version") else None
50+
suppress_invoke_agent_input = kwargs.get("suppress_invoke_agent_input", False)
5051

5152
# Get the configured Microsoft Agent 365 Tracer
5253
try:
@@ -64,7 +65,9 @@ def _instrument(self, **kwargs: Any) -> None:
6465

6566
agent365_tracer = cast(Tracer, tracer)
6667

67-
set_trace_processors([OpenAIAgentsTraceProcessor(agent365_tracer)])
68+
set_trace_processors(
69+
[OpenAIAgentsTraceProcessor(agent365_tracer, suppress_invoke_agent_input)]
70+
)
6871

6972
def _uninstrument(self, **kwargs: Any) -> None:
7073
pass

libraries/microsoft-agents-a365-observability-extensions-openai/microsoft_agents_a365/observability/extensions/openai/trace_processor.py

Lines changed: 38 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -69,8 +69,9 @@
6969
class OpenAIAgentsTraceProcessor(TracingProcessor):
7070
_MAX_HANDOFFS_IN_FLIGHT = 1000
7171

72-
def __init__(self, tracer: Tracer) -> None:
72+
def __init__(self, tracer: Tracer, suppress_invoke_agent_input: bool = False) -> None:
7373
self._tracer = tracer
74+
self._suppress_invoke_agent_input = suppress_invoke_agent_input
7475
self._root_spans: dict[str, OtelSpan] = {}
7576
self._otel_spans: dict[str, OtelSpan] = {}
7677
self._tokens: dict[str, object] = {}
@@ -79,6 +80,8 @@ def __init__(self, tracer: Tracer) -> None:
7980
# Use an OrderedDict and _MAX_HANDOFFS_IN_FLIGHT to cap the size of the dict
8081
# in case there are large numbers of orphaned handoffs
8182
self._reverse_handoffs_dict: OrderedDict[str, str] = OrderedDict()
83+
# Track active agent spans per trace to determine if we're in an InvokeAgent scope
84+
self._active_agent_spans: dict[str, set[str]] = {}
8285

8386
# helper
8487
def _stamp_custom_parent(self, otel_span: OtelSpan, trace_id: str) -> None:
@@ -89,6 +92,11 @@ def _stamp_custom_parent(self, otel_span: OtelSpan, trace_id: str) -> None:
8992
pid_hex = "0x" + ot_trace.format_span_id(sc.span_id)
9093
otel_span.set_attribute(CUSTOM_PARENT_SPAN_ID_KEY, pid_hex)
9194

95+
def _is_in_invoke_agent_scope(self, trace_id: str) -> bool:
96+
"""Check if we're currently inside an InvokeAgent scope for the given trace."""
97+
agent_spans = self._active_agent_spans.get(trace_id, set())
98+
return len(agent_spans) > 0
99+
92100
def on_trace_start(self, trace: Trace) -> None:
93101
"""Called when a trace is started.
94102
@@ -134,6 +142,12 @@ def on_span_start(self, span: Span[Any]) -> None:
134142
self._otel_spans[span.span_id] = otel_span
135143
self._tokens[span.span_id] = attach(set_span_in_context(otel_span))
136144

145+
# Track agent spans for InvokeAgent scope detection
146+
if isinstance(span.span_data, AgentSpanData):
147+
if span.trace_id not in self._active_agent_spans:
148+
self._active_agent_spans[span.trace_id] = set()
149+
self._active_agent_spans[span.trace_id].add(span.span_id)
150+
137151
def on_span_end(self, span: Span[Any]) -> None:
138152
"""Called when a span is finished. Should not block or raise exceptions.
139153
@@ -154,7 +168,11 @@ def on_span_end(self, span: Span[Any]) -> None:
154168
otel_span.set_attribute(GEN_AI_OUTPUT_MESSAGES_KEY, response.model_dump_json())
155169
for k, v in get_attributes_from_response(response):
156170
otel_span.set_attribute(k, v)
157-
if hasattr(data, "input") and (input := data.input):
171+
# Only record input messages if not suppressing or not in InvokeAgent scope
172+
should_suppress = self._suppress_invoke_agent_input and self._is_in_invoke_agent_scope(
173+
span.trace_id
174+
)
175+
if not should_suppress and hasattr(data, "input") and (input := data.input):
158176
if isinstance(input, str):
159177
otel_span.set_attribute(GEN_AI_INPUT_MESSAGES_KEY, input)
160178
elif isinstance(input, list):
@@ -164,8 +182,18 @@ def on_span_end(self, span: Span[Any]) -> None:
164182
elif TYPE_CHECKING:
165183
assert_never(input)
166184
elif isinstance(data, GenerationSpanData):
167-
for k, v in get_attributes_from_generation_span_data(data):
168-
otel_span.set_attribute(k, v)
185+
# Only record input messages if not suppressing or not in InvokeAgent scope
186+
should_suppress = self._suppress_invoke_agent_input and self._is_in_invoke_agent_scope(
187+
span.trace_id
188+
)
189+
if not should_suppress:
190+
for k, v in get_attributes_from_generation_span_data(data):
191+
otel_span.set_attribute(k, v)
192+
else:
193+
# Still set attributes other than input messages
194+
for k, v in get_attributes_from_generation_span_data(data):
195+
if k != GEN_AI_INPUT_MESSAGES_KEY:
196+
otel_span.set_attribute(k, v)
169197
self._stamp_custom_parent(otel_span, span.trace_id)
170198
otel_span.update_name(
171199
f"{otel_span.attributes[GEN_AI_OPERATION_NAME_KEY]} {otel_span.attributes[GEN_AI_REQUEST_MODEL_KEY]}"
@@ -194,6 +222,12 @@ def on_span_end(self, span: Span[Any]) -> None:
194222
otel_span.set_attribute(GEN_AI_GRAPH_NODE_PARENT_ID, parent_node)
195223
otel_span.update_name(f"{INVOKE_AGENT_OPERATION_NAME} {get_span_name(span)}")
196224

225+
# Clean up agent span tracking
226+
if span.trace_id in self._active_agent_spans:
227+
self._active_agent_spans[span.trace_id].discard(span.span_id)
228+
if not self._active_agent_spans[span.trace_id]:
229+
del self._active_agent_spans[span.trace_id]
230+
197231
end_time: int | None = None
198232
if span.ended_at:
199233
try:
Lines changed: 283 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,283 @@
1+
# Copyright (c) Microsoft. All rights reserved.
2+
3+
import unittest
4+
from datetime import datetime
5+
from unittest.mock import Mock
6+
7+
from agents.tracing import Span
8+
from agents.tracing.span_data import AgentSpanData, GenerationSpanData, ResponseSpanData
9+
from microsoft_agents_a365.observability.core import configure, get_tracer
10+
from microsoft_agents_a365.observability.core.constants import GEN_AI_INPUT_MESSAGES_KEY
11+
from microsoft_agents_a365.observability.extensions.openai.trace_processor import (
12+
OpenAIAgentsTraceProcessor,
13+
)
14+
from openai.types.responses import Response
15+
16+
17+
class TestPromptSuppression(unittest.TestCase):
18+
"""Unit tests for prompt suppression functionality in OpenAIAgentsTraceProcessor."""
19+
20+
@classmethod
21+
def setUpClass(cls):
22+
"""Set up test environment once for all tests."""
23+
configure(
24+
service_name="test-service-prompt-suppression",
25+
service_namespace="test-namespace-prompt-suppression",
26+
)
27+
28+
def setUp(self):
29+
"""Set up each test with a fresh processor and mock tracer."""
30+
self.tracer = get_tracer()
31+
self.mock_otel_span = Mock()
32+
self.mock_otel_span.attributes = {}
33+
self.mock_otel_span.get_span_context.return_value = Mock(
34+
trace_id="test-trace-id", span_id="test-span-id"
35+
)
36+
37+
# Track attributes set on the span
38+
def set_attribute_side_effect(key, value):
39+
self.mock_otel_span.attributes[key] = value
40+
41+
self.mock_otel_span.set_attribute = Mock(side_effect=set_attribute_side_effect)
42+
self.mock_otel_span.update_name = Mock()
43+
self.mock_otel_span.set_status = Mock()
44+
self.mock_otel_span.end = Mock()
45+
46+
# Mock the tracer's start_span method
47+
self.original_start_span = self.tracer.start_span
48+
self.tracer.start_span = Mock(return_value=self.mock_otel_span)
49+
50+
def tearDown(self):
51+
"""Clean up after each test."""
52+
self.tracer.start_span = self.original_start_span
53+
54+
def test_does_not_record_input_messages_when_suppression_enabled_in_agent_scope(self):
55+
"""Test that input messages are not recorded when suppression is enabled and in agent scope."""
56+
processor = OpenAIAgentsTraceProcessor(self.tracer, suppress_invoke_agent_input=True)
57+
58+
trace_id = "trace-suppress"
59+
now = datetime.now().isoformat()
60+
61+
# Start an agent span to create InvokeAgent scope
62+
agent_span = Mock(spec=Span)
63+
agent_span.span_id = "agent-span"
64+
agent_span.trace_id = trace_id
65+
agent_span.parent_id = None
66+
agent_span.started_at = now
67+
agent_span.ended_at = None
68+
agent_span.span_data = AgentSpanData(name="TestAgent")
69+
70+
processor.on_span_start(agent_span)
71+
72+
# Now create a generation span with input (proper format - list of message dicts)
73+
gen_span = Mock(spec=Span)
74+
gen_span.span_id = "gen-span"
75+
gen_span.trace_id = trace_id
76+
gen_span.parent_id = "agent-span"
77+
gen_span.started_at = now
78+
gen_span.ended_at = now
79+
gen_span.span_data = GenerationSpanData(
80+
model="gpt-4",
81+
input=[{"role": "user", "content": "Hello prompt"}]
82+
)
83+
84+
processor.on_span_start(gen_span)
85+
processor.on_span_end(gen_span)
86+
87+
# Verify that set_attribute was called but NOT with GEN_AI_INPUT_MESSAGES_KEY
88+
attribute_keys = [call[0][0] for call in self.mock_otel_span.set_attribute.call_args_list]
89+
self.assertNotIn(
90+
GEN_AI_INPUT_MESSAGES_KEY,
91+
attribute_keys,
92+
"GEN_AI_INPUT_MESSAGES_KEY should not be set when suppression is enabled",
93+
)
94+
95+
def test_records_input_messages_when_suppression_disabled(self):
96+
"""Test that input messages are recorded when suppression is disabled (default)."""
97+
processor = OpenAIAgentsTraceProcessor(self.tracer, suppress_invoke_agent_input=False)
98+
99+
trace_id = "trace-allow"
100+
now = datetime.now().isoformat()
101+
102+
# Start an agent span
103+
agent_span = Mock(spec=Span)
104+
agent_span.span_id = "agent-span-2"
105+
agent_span.trace_id = trace_id
106+
agent_span.parent_id = None
107+
agent_span.started_at = now
108+
agent_span.ended_at = None
109+
agent_span.span_data = AgentSpanData(name="TestAgent")
110+
111+
processor.on_span_start(agent_span)
112+
113+
# Create a generation span with input (proper format - list of message dicts)
114+
gen_span = Mock(spec=Span)
115+
gen_span.span_id = "gen-span-2"
116+
gen_span.trace_id = trace_id
117+
gen_span.parent_id = "agent-span-2"
118+
gen_span.started_at = now
119+
gen_span.ended_at = now
120+
gen_span.span_data = GenerationSpanData(
121+
model="gpt-4",
122+
input=[{"role": "user", "content": "Hello prompt"}]
123+
)
124+
125+
processor.on_span_start(gen_span)
126+
processor.on_span_end(gen_span)
127+
128+
# Verify that set_attribute was called with GEN_AI_INPUT_MESSAGES_KEY
129+
attribute_keys = [call[0][0] for call in self.mock_otel_span.set_attribute.call_args_list]
130+
self.assertIn(
131+
GEN_AI_INPUT_MESSAGES_KEY,
132+
attribute_keys,
133+
"GEN_AI_INPUT_MESSAGES_KEY should be set when suppression is disabled",
134+
)
135+
136+
def test_suppresses_input_on_response_spans_when_enabled(self):
137+
"""Test that input is suppressed on response spans when suppression is enabled."""
138+
processor = OpenAIAgentsTraceProcessor(self.tracer, suppress_invoke_agent_input=True)
139+
140+
trace_id = "trace-resp"
141+
now = datetime.now().isoformat()
142+
143+
# Start an agent span
144+
agent_span = Mock(spec=Span)
145+
agent_span.span_id = "agent-span-3"
146+
agent_span.trace_id = trace_id
147+
agent_span.parent_id = None
148+
agent_span.started_at = now
149+
agent_span.ended_at = None
150+
agent_span.span_data = AgentSpanData(name="TestAgent")
151+
152+
processor.on_span_start(agent_span)
153+
154+
# Create a response span with input
155+
resp_span = Mock(spec=Span)
156+
resp_span.span_id = "resp-span"
157+
resp_span.trace_id = trace_id
158+
resp_span.parent_id = "agent-span-3"
159+
resp_span.started_at = now
160+
resp_span.ended_at = now
161+
162+
# Create mock response data with all required attributes
163+
mock_response = Mock(spec=Response)
164+
mock_response.model_dump_json.return_value = '{"output": "test"}'
165+
mock_response.tools = None
166+
mock_response.usage = None
167+
mock_response.output = None
168+
mock_response.instructions = None
169+
mock_response.model = "gpt-4"
170+
mock_response.model_dump.return_value = {}
171+
172+
resp_span.span_data = Mock(spec=ResponseSpanData)
173+
resp_span.span_data.response = mock_response
174+
resp_span.span_data.input = "Prompt text"
175+
176+
processor.on_span_start(resp_span)
177+
processor.on_span_end(resp_span)
178+
179+
# Verify that set_attribute was called but NOT with GEN_AI_INPUT_MESSAGES_KEY for input
180+
attribute_keys = [call[0][0] for call in self.mock_otel_span.set_attribute.call_args_list]
181+
self.assertNotIn(
182+
GEN_AI_INPUT_MESSAGES_KEY,
183+
attribute_keys,
184+
"GEN_AI_INPUT_MESSAGES_KEY should not be set for response span when suppression is enabled",
185+
)
186+
187+
def test_records_input_outside_agent_scope_even_when_suppression_enabled(self):
188+
"""Test that input messages are recorded outside agent scope even when suppression is enabled."""
189+
processor = OpenAIAgentsTraceProcessor(self.tracer, suppress_invoke_agent_input=True)
190+
191+
trace_id = "trace-outside"
192+
now = datetime.now().isoformat()
193+
194+
# Create a generation span WITHOUT an agent span (outside InvokeAgent scope)
195+
gen_span = Mock(spec=Span)
196+
gen_span.span_id = "gen-span-outside"
197+
gen_span.trace_id = trace_id
198+
gen_span.parent_id = None
199+
gen_span.started_at = now
200+
gen_span.ended_at = now
201+
gen_span.span_data = GenerationSpanData(
202+
model="gpt-4",
203+
input=[{"role": "user", "content": "Hello prompt"}]
204+
)
205+
206+
processor.on_span_start(gen_span)
207+
processor.on_span_end(gen_span)
208+
209+
# Verify that set_attribute WAS called with GEN_AI_INPUT_MESSAGES_KEY
210+
# because we're not in an InvokeAgent scope
211+
attribute_keys = [call[0][0] for call in self.mock_otel_span.set_attribute.call_args_list]
212+
self.assertIn(
213+
GEN_AI_INPUT_MESSAGES_KEY,
214+
attribute_keys,
215+
"GEN_AI_INPUT_MESSAGES_KEY should be set when outside InvokeAgent scope",
216+
)
217+
218+
def test_default_suppression_is_false(self):
219+
"""Test that the default value for suppress_invoke_agent_input is False."""
220+
processor = OpenAIAgentsTraceProcessor(self.tracer)
221+
222+
self.assertFalse(
223+
processor._suppress_invoke_agent_input,
224+
"Default value for suppress_invoke_agent_input should be False",
225+
)
226+
227+
def test_agent_span_tracking_cleanup(self):
228+
"""Test that agent span tracking is properly cleaned up when spans end."""
229+
processor = OpenAIAgentsTraceProcessor(self.tracer, suppress_invoke_agent_input=True)
230+
231+
trace_id = "trace-cleanup"
232+
now = datetime.now().isoformat()
233+
234+
# Start an agent span
235+
agent_span = Mock(spec=Span)
236+
agent_span.span_id = "agent-span-cleanup"
237+
agent_span.trace_id = trace_id
238+
agent_span.parent_id = None
239+
agent_span.started_at = now
240+
agent_span.ended_at = now
241+
agent_span.span_data = AgentSpanData(name="TestAgent")
242+
243+
processor.on_span_start(agent_span)
244+
245+
# Verify the span is tracked
246+
self.assertIn(trace_id, processor._active_agent_spans)
247+
self.assertIn(agent_span.span_id, processor._active_agent_spans[trace_id])
248+
249+
# End the agent span
250+
processor.on_span_end(agent_span)
251+
252+
# Verify the tracking is cleaned up
253+
self.assertNotIn(trace_id, processor._active_agent_spans)
254+
255+
256+
def run_tests():
257+
"""Run all prompt suppression tests."""
258+
print("🧪 Running prompt suppression tests...")
259+
print("=" * 80)
260+
261+
loader = unittest.TestLoader()
262+
suite = loader.loadTestsFromTestCase(TestPromptSuppression)
263+
264+
runner = unittest.TextTestRunner(verbosity=2)
265+
result = runner.run(suite)
266+
267+
print("\n" + "=" * 80)
268+
print("🏁 Test Summary:")
269+
print(f"Tests run: {result.testsRun}")
270+
print(f"Failures: {len(result.failures)}")
271+
print(f"Errors: {len(result.errors)}")
272+
273+
if result.wasSuccessful():
274+
print("🎉 All tests passed!")
275+
return True
276+
else:
277+
print("🔧 Some tests failed. Check output above.")
278+
return False
279+
280+
281+
if __name__ == "__main__":
282+
success = run_tests()
283+
exit(0 if success else 1)

0 commit comments

Comments
 (0)