Skip to content

Commit f02e68c

Browse files
CopilotnikhilNava
andcommitted
feat: add bounded collections for LangChain tracer and OutputScope
- Convert LangChain _spans_by_run from unbounded DictWithLock to bounded OrderedDict with _MAX_TRACKED_RUNS=10000 cap - Add _cap_ordered_dict helper for FIFO eviction (matching OpenAI pattern) - Add thread-safe lock usage for _spans_by_run in error handlers - Add _MAX_OUTPUT_MESSAGES=5000 cap for OutputScope._output_messages - Add unit tests for both bounded collections Co-authored-by: nikhilNava <211831449+nikhilNava@users.noreply.github.com>
1 parent c08b24f commit f02e68c

4 files changed

Lines changed: 241 additions & 7 deletions

File tree

‎libraries/microsoft-agents-a365-observability-core/microsoft_agents_a365/observability/core/spans_scopes/output_scope.py‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,8 @@
1616
class OutputScope(OpenTelemetryScope):
1717
"""Provides OpenTelemetry tracing scope for output messages."""
1818

19+
_MAX_OUTPUT_MESSAGES = 5000
20+
1921
@staticmethod
2022
def start(
2123
agent_details: AgentDetails,
@@ -82,9 +84,12 @@ def record_output_messages(self, messages: list[str]) -> None:
8284
"""Records the output messages for telemetry tracking.
8385
8486
Appends the provided messages to the accumulated output messages list.
87+
The list is capped at _MAX_OUTPUT_MESSAGES to prevent unbounded memory growth.
8588
8689
Args:
8790
messages: List of output messages to append
8891
"""
8992
self._output_messages.extend(messages)
93+
if len(self._output_messages) > self._MAX_OUTPUT_MESSAGES:
94+
self._output_messages = self._output_messages[-self._MAX_OUTPUT_MESSAGES :]
9095
self.set_tag_maybe(GEN_AI_OUTPUT_MESSAGES_KEY, safe_json_dumps(self._output_messages))

‎libraries/microsoft-agents-a365-observability-extensions-langchain/microsoft_agents_a365/observability/extensions/langchain/tracer.py‎

Lines changed: 27 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33

44
import logging
55
import re
6+
from collections import OrderedDict
67
from collections.abc import Iterator
78
from itertools import chain
89
from threading import RLock
@@ -69,6 +70,8 @@
6970

7071

7172
class CustomLangChainTracer(BaseTracer):
73+
_MAX_TRACKED_RUNS = 10000
74+
7275
__slots__ = (
7376
"_tracer",
7477
"_separate_trace_from_runtime_context",
@@ -98,11 +101,18 @@ def __init__(
98101
self.run_map = DictWithLock[str, Run](self.run_map)
99102
self._tracer = tracer
100103
self._separate_trace_from_runtime_context = separate_trace_from_runtime_context
101-
self._spans_by_run: dict[UUID, Span] = DictWithLock[UUID, Span]()
104+
self._spans_by_run: OrderedDict[UUID, Span] = OrderedDict()
102105
self._lock = RLock() # handlers may be run in a thread by langchain
103106

104107
def get_span(self, run_id: UUID) -> Span | None:
105-
return self._spans_by_run.get(run_id)
108+
with self._lock:
109+
return self._spans_by_run.get(run_id)
110+
111+
@staticmethod
112+
def _cap_ordered_dict(d: OrderedDict, max_size: int) -> None:
113+
"""Evict oldest entries from an OrderedDict to stay within max_size."""
114+
while len(d) > max_size:
115+
d.popitem(last=False)
106116

107117
def _start_trace(self, run: Run) -> None:
108118
self.run_map[str(run.id)] = run
@@ -142,12 +152,14 @@ def _start_trace(self, run: Run) -> None:
142152
# token = context_api.attach(context)
143153
with self._lock:
144154
self._spans_by_run[run.id] = span
155+
self._cap_ordered_dict(self._spans_by_run, self._MAX_TRACKED_RUNS)
145156

146157
def _end_trace(self, run: Run) -> None:
147158
self.run_map.pop(str(run.id), None)
148159
if context_api.get_value(_SUPPRESS_INSTRUMENTATION_KEY):
149160
return
150-
span = self._spans_by_run.pop(run.id, None)
161+
with self._lock:
162+
span = self._spans_by_run.pop(run.id, None)
151163
if span:
152164
try:
153165
_update_span(span, run)
@@ -162,24 +174,32 @@ def _persist_run(self, run: Run) -> None:
162174
pass
163175

164176
def on_llm_error(self, error: BaseException, *args: Any, run_id: UUID, **kwargs: Any) -> Run:
165-
if span := self._spans_by_run.get(run_id):
177+
with self._lock:
178+
span = self._spans_by_run.get(run_id)
179+
if span:
166180
record_exception(span, error)
167181
return super().on_llm_error(error, *args, run_id=run_id, **kwargs)
168182

169183
def on_chain_error(self, error: BaseException, *args: Any, run_id: UUID, **kwargs: Any) -> Run:
170-
if span := self._spans_by_run.get(run_id):
184+
with self._lock:
185+
span = self._spans_by_run.get(run_id)
186+
if span:
171187
record_exception(span, error)
172188
return super().on_chain_error(error, *args, run_id=run_id, **kwargs)
173189

174190
def on_retriever_error(
175191
self, error: BaseException, *args: Any, run_id: UUID, **kwargs: Any
176192
) -> Run:
177-
if span := self._spans_by_run.get(run_id):
193+
with self._lock:
194+
span = self._spans_by_run.get(run_id)
195+
if span:
178196
record_exception(span, error)
179197
return super().on_retriever_error(error, *args, run_id=run_id, **kwargs)
180198

181199
def on_tool_error(self, error: BaseException, *args: Any, run_id: UUID, **kwargs: Any) -> Run:
182-
if span := self._spans_by_run.get(run_id):
200+
with self._lock:
201+
span = self._spans_by_run.get(run_id)
202+
if span:
183203
record_exception(span, error)
184204
return super().on_tool_error(error, *args, run_id=run_id, **kwargs)
185205

Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT License.
3+
4+
"""Tests for bounded output messages in OutputScope."""
5+
6+
import unittest
7+
from unittest.mock import MagicMock, patch
8+
9+
from microsoft_agents_a365.observability.core.spans_scopes.output_scope import OutputScope
10+
11+
12+
class TestOutputScopeBounded(unittest.TestCase):
13+
"""Tests that OutputScope._output_messages list is properly bounded."""
14+
15+
def _make_scope(self, initial_messages: list[str] | None = None) -> OutputScope:
16+
"""Create an OutputScope with mocked dependencies."""
17+
agent_details = MagicMock()
18+
agent_details.agent_id = "test-agent"
19+
agent_details.agent_name = "Test Agent"
20+
agent_details.agent_description = None
21+
agent_details.platform_id = None
22+
agent_details.conversation_id = None
23+
agent_details.icon_uri = None
24+
agent_details.agent_auid = None
25+
agent_details.agent_upn = None
26+
agent_details.agent_blueprint_id = None
27+
28+
tenant_details = MagicMock()
29+
tenant_details.tenant_id = "test-tenant"
30+
31+
response = MagicMock()
32+
response.messages = initial_messages or ["hello"]
33+
34+
with patch.object(OutputScope, "__init__", lambda self, *a, **kw: None):
35+
scope = OutputScope.__new__(OutputScope)
36+
scope._output_messages = list(response.messages)
37+
scope.set_tag_maybe = MagicMock()
38+
39+
return scope
40+
41+
def test_max_output_messages_default(self):
42+
"""Default _MAX_OUTPUT_MESSAGES should be 5000."""
43+
self.assertEqual(OutputScope._MAX_OUTPUT_MESSAGES, 5000)
44+
45+
def test_record_output_messages_within_limit(self):
46+
"""Messages under the limit should not be truncated."""
47+
scope = self._make_scope(["initial"])
48+
scope.record_output_messages(["msg1", "msg2", "msg3"])
49+
self.assertEqual(len(scope._output_messages), 4)
50+
self.assertEqual(scope._output_messages, ["initial", "msg1", "msg2", "msg3"])
51+
52+
def test_record_output_messages_exceeds_limit(self):
53+
"""Messages exceeding the limit should be truncated to keep newest."""
54+
scope = self._make_scope([])
55+
original_max = OutputScope._MAX_OUTPUT_MESSAGES
56+
try:
57+
OutputScope._MAX_OUTPUT_MESSAGES = 10
58+
59+
# Add 15 messages
60+
scope.record_output_messages([f"msg_{i}" for i in range(15)])
61+
62+
# Should be capped at 10 (keeping the newest)
63+
self.assertEqual(len(scope._output_messages), 10)
64+
# Oldest 5 should be gone, newest 10 should remain
65+
self.assertEqual(scope._output_messages[0], "msg_5")
66+
self.assertEqual(scope._output_messages[-1], "msg_14")
67+
finally:
68+
OutputScope._MAX_OUTPUT_MESSAGES = original_max
69+
70+
def test_record_output_messages_multiple_calls_capped(self):
71+
"""Multiple calls to record_output_messages should stay bounded."""
72+
scope = self._make_scope([])
73+
original_max = OutputScope._MAX_OUTPUT_MESSAGES
74+
try:
75+
OutputScope._MAX_OUTPUT_MESSAGES = 5
76+
77+
for batch in range(4):
78+
scope.record_output_messages([f"batch{batch}_msg{i}" for i in range(3)])
79+
80+
# Total of 12 messages added in 4 batches, should be capped at 5
81+
self.assertLessEqual(len(scope._output_messages), 5)
82+
# Latest messages should be from the last batches
83+
self.assertIn("batch3_msg2", scope._output_messages)
84+
finally:
85+
OutputScope._MAX_OUTPUT_MESSAGES = original_max
86+
87+
def test_record_output_messages_exactly_at_limit(self):
88+
"""Messages exactly at the limit should not be truncated."""
89+
scope = self._make_scope([])
90+
original_max = OutputScope._MAX_OUTPUT_MESSAGES
91+
try:
92+
OutputScope._MAX_OUTPUT_MESSAGES = 5
93+
scope.record_output_messages([f"msg_{i}" for i in range(5)])
94+
self.assertEqual(len(scope._output_messages), 5)
95+
finally:
96+
OutputScope._MAX_OUTPUT_MESSAGES = original_max
97+
98+
99+
if __name__ == "__main__":
100+
unittest.main()
Lines changed: 109 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,109 @@
1+
# Copyright (c) Microsoft Corporation.
2+
# Licensed under the MIT License.
3+
4+
"""Tests for bounded collections in the LangChain tracer."""
5+
6+
import unittest
7+
from collections import OrderedDict
8+
from unittest.mock import MagicMock, patch
9+
from uuid import uuid4
10+
11+
from microsoft_agents_a365.observability.extensions.langchain.tracer import (
12+
CustomLangChainTracer,
13+
)
14+
15+
16+
class TestLangChainTracerBounded(unittest.TestCase):
17+
"""Tests that LangChain tracer collections are properly bounded."""
18+
19+
def _make_tracer(self) -> CustomLangChainTracer:
20+
"""Create a tracer with a mock OTel tracer."""
21+
mock_otel_tracer = MagicMock()
22+
mock_span = MagicMock()
23+
mock_otel_tracer.start_span.return_value = mock_span
24+
return CustomLangChainTracer(
25+
tracer=mock_otel_tracer,
26+
separate_trace_from_runtime_context=True,
27+
)
28+
29+
def test_spans_by_run_is_ordered_dict(self):
30+
"""_spans_by_run should be an OrderedDict for bounded eviction."""
31+
tracer = self._make_tracer()
32+
self.assertIsInstance(tracer._spans_by_run, OrderedDict)
33+
34+
def test_cap_ordered_dict_evicts_oldest(self):
35+
"""_cap_ordered_dict should evict oldest entries (FIFO)."""
36+
d: OrderedDict[str, int] = OrderedDict()
37+
for i in range(15):
38+
d[f"key_{i}"] = i
39+
CustomLangChainTracer._cap_ordered_dict(d, 10)
40+
41+
self.assertEqual(len(d), 10)
42+
# oldest 5 should be gone
43+
for i in range(5):
44+
self.assertNotIn(f"key_{i}", d)
45+
# newest 10 should remain
46+
for i in range(5, 15):
47+
self.assertIn(f"key_{i}", d)
48+
self.assertEqual(d[f"key_{i}"], i)
49+
50+
def test_cap_ordered_dict_noop_when_under_limit(self):
51+
"""_cap_ordered_dict should be a no-op when size is under limit."""
52+
d: OrderedDict[str, int] = OrderedDict()
53+
for i in range(5):
54+
d[f"key_{i}"] = i
55+
CustomLangChainTracer._cap_ordered_dict(d, 10)
56+
self.assertEqual(len(d), 5)
57+
58+
def test_spans_by_run_bounded_on_start_trace(self):
59+
"""_spans_by_run should be bounded when _start_trace adds entries."""
60+
tracer = self._make_tracer()
61+
# Use a small cap for testing
62+
original_max = CustomLangChainTracer._MAX_TRACKED_RUNS
63+
try:
64+
CustomLangChainTracer._MAX_TRACKED_RUNS = 5
65+
66+
# Add more runs than the cap
67+
for i in range(10):
68+
run = MagicMock()
69+
run.id = uuid4()
70+
run.parent_run_id = None
71+
run.run_type = "llm"
72+
run.name = f"test_run_{i}"
73+
run.start_time = MagicMock()
74+
75+
with patch(
76+
"microsoft_agents_a365.observability.extensions.langchain.tracer"
77+
".context_api.get_value",
78+
return_value=None,
79+
):
80+
tracer._start_trace(run)
81+
82+
# Should be capped at 5
83+
self.assertLessEqual(len(tracer._spans_by_run), 5)
84+
finally:
85+
CustomLangChainTracer._MAX_TRACKED_RUNS = original_max
86+
87+
def test_get_span_returns_none_for_missing(self):
88+
"""get_span should return None for non-existent run_id."""
89+
tracer = self._make_tracer()
90+
result = tracer.get_span(uuid4())
91+
self.assertIsNone(result)
92+
93+
def test_get_span_returns_span_for_existing(self):
94+
"""get_span should return the span for existing run_id."""
95+
tracer = self._make_tracer()
96+
run_id = uuid4()
97+
mock_span = MagicMock()
98+
with tracer._lock:
99+
tracer._spans_by_run[run_id] = mock_span
100+
result = tracer.get_span(run_id)
101+
self.assertEqual(result, mock_span)
102+
103+
def test_max_tracked_runs_default(self):
104+
"""Default _MAX_TRACKED_RUNS should be 10000."""
105+
self.assertEqual(CustomLangChainTracer._MAX_TRACKED_RUNS, 10000)
106+
107+
108+
if __name__ == "__main__":
109+
unittest.main()

0 commit comments

Comments
 (0)