Skip to content

Commit 8d89bfb

Browse files
alvinkam2001stainless-app[bot]
authored andcommitted
batch SGP span upserts (#331)
1 parent 23fb7e4 commit 8d89bfb

6 files changed

Lines changed: 352 additions & 46 deletions

File tree

‎src/agentex/lib/core/tracing/processors/sgp_tracing_processor.py‎

Lines changed: 48 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
1+
from __future__ import annotations
2+
13
from typing import override
24

35
import scale_gp_beta.lib.tracing as tracing
@@ -125,48 +127,64 @@ def _add_source_to_span(self, span: Span) -> None:
125127

126128
@override
127129
async def on_span_start(self, span: Span) -> None:
128-
self._add_source_to_span(span)
129-
sgp_span = create_span(
130-
name=span.name,
131-
span_type=_get_span_type(span),
132-
span_id=span.id,
133-
parent_id=span.parent_id,
134-
trace_id=span.trace_id,
135-
input=span.input,
136-
output=span.output,
137-
metadata=span.data,
138-
)
139-
sgp_span.start_time = span.start_time.isoformat() # type: ignore[union-attr]
130+
await self.on_spans_start([span])
131+
132+
@override
133+
async def on_span_end(self, span: Span) -> None:
134+
await self.on_spans_end([span])
135+
136+
@override
137+
async def on_spans_start(self, spans: list[Span]) -> None:
138+
if not spans:
139+
return
140+
141+
sgp_spans: list[SGPSpan] = []
142+
for span in spans:
143+
self._add_source_to_span(span)
144+
sgp_span = create_span(
145+
name=span.name,
146+
span_type=_get_span_type(span),
147+
span_id=span.id,
148+
parent_id=span.parent_id,
149+
trace_id=span.trace_id,
150+
input=span.input,
151+
output=span.output,
152+
metadata=span.data,
153+
)
154+
sgp_span.start_time = span.start_time.isoformat() # type: ignore[union-attr]
155+
self._spans[span.id] = sgp_span
156+
sgp_spans.append(sgp_span)
140157

141158
if self.disabled:
142159
logger.warning("SGP is disabled, skipping span upsert")
143160
return
144-
# TODO(AGX1-198): Batch multiple spans into a single upsert_batch call
145-
# instead of one span per HTTP request.
146-
# https://linear.app/scale-epd/issue/AGX1-198/actually-use-sgp-batching-for-spans
147161
await self.sgp_async_client.spans.upsert_batch( # type: ignore[union-attr]
148-
items=[sgp_span.to_request_params()]
162+
items=[s.to_request_params() for s in sgp_spans]
149163
)
150164

151-
self._spans[span.id] = sgp_span
152-
153165
@override
154-
async def on_span_end(self, span: Span) -> None:
155-
sgp_span = self._spans.pop(span.id, None)
156-
if sgp_span is None:
157-
logger.warning(f"Span {span.id} not found in stored spans, skipping span end")
166+
async def on_spans_end(self, spans: list[Span]) -> None:
167+
if not spans:
158168
return
159169

160-
self._add_source_to_span(span)
161-
sgp_span.input = span.input # type: ignore[assignment]
162-
sgp_span.output = span.output # type: ignore[assignment]
163-
sgp_span.metadata = span.data # type: ignore[assignment]
164-
sgp_span.end_time = span.end_time.isoformat() # type: ignore[union-attr]
165-
166-
if self.disabled:
170+
to_upsert: list[SGPSpan] = []
171+
for span in spans:
172+
sgp_span = self._spans.pop(span.id, None)
173+
if sgp_span is None:
174+
logger.warning(f"Span {span.id} not found in stored spans, skipping span end")
175+
continue
176+
177+
self._add_source_to_span(span)
178+
sgp_span.input = span.input # type: ignore[assignment]
179+
sgp_span.output = span.output # type: ignore[assignment]
180+
sgp_span.metadata = span.data # type: ignore[assignment]
181+
sgp_span.end_time = span.end_time.isoformat() # type: ignore[union-attr]
182+
to_upsert.append(sgp_span)
183+
184+
if self.disabled or not to_upsert:
167185
return
168186
await self.sgp_async_client.spans.upsert_batch( # type: ignore[union-attr]
169-
items=[sgp_span.to_request_params()]
187+
items=[s.to_request_params() for s in to_upsert]
170188
)
171189

172190
@override

‎src/agentex/lib/core/tracing/processors/tracing_processor_interface.py‎

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,13 @@
1+
from __future__ import annotations
2+
3+
import asyncio
14
from abc import ABC, abstractmethod
25

36
from agentex.types.span import Span
47
from agentex.lib.types.tracing import TracingProcessorConfig
8+
from agentex.lib.utils.logging import make_logger
9+
10+
logger = make_logger(__name__)
511

612

713
class SyncTracingProcessor(ABC):
@@ -35,6 +41,43 @@ async def on_span_start(self, span: Span) -> None:
3541
async def on_span_end(self, span: Span) -> None:
3642
pass
3743

44+
async def on_spans_start(self, spans: list[Span]) -> None:
45+
"""Batched variant of on_span_start.
46+
47+
Default fallback fans out to the single-span method in parallel so
48+
existing processors keep working unchanged. Processors that support
49+
real batching (e.g. sending all spans in one HTTP call) should
50+
override this to avoid the per-span round trip.
51+
52+
Per-span exceptions are captured and logged individually so that one
53+
failing span does not prevent the others from being processed.
54+
"""
55+
results = await asyncio.gather(
56+
*(self.on_span_start(s) for s in spans), return_exceptions=True
57+
)
58+
for span, result in zip(spans, results):
59+
if isinstance(result, Exception):
60+
logger.error(
61+
"Tracing processor %s failed on_span_start for span %s",
62+
type(self).__name__,
63+
span.id,
64+
exc_info=result,
65+
)
66+
67+
async def on_spans_end(self, spans: list[Span]) -> None:
68+
"""Batched variant of on_span_end. See on_spans_start for details."""
69+
results = await asyncio.gather(
70+
*(self.on_span_end(s) for s in spans), return_exceptions=True
71+
)
72+
for span, result in zip(spans, results):
73+
if isinstance(result, Exception):
74+
logger.error(
75+
"Tracing processor %s failed on_span_end for span %s",
76+
type(self).__name__,
77+
span.id,
78+
exc_info=result,
79+
)
80+
3881
@abstractmethod
3982
async def shutdown(self) -> None:
4083
pass

‎src/agentex/lib/core/tracing/span_queue.py‎

Lines changed: 27 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -95,29 +95,40 @@ async def _drain_loop(self) -> None:
9595

9696
@staticmethod
9797
async def _process_items(items: list[_SpanQueueItem]) -> None:
98-
"""Process a list of span events concurrently."""
98+
"""Dispatch a batch of same-event-type items to each processor in one call.
9999
100-
async def _handle(item: _SpanQueueItem) -> None:
100+
Groups spans by processor so each processor sees its full slice of the
101+
drain batch at once. Processors that override the batched methods can
102+
then send a single HTTP request per drain cycle instead of N.
103+
"""
104+
if not items:
105+
return
106+
107+
event_type = items[0].event_type
108+
assert all(i.event_type == event_type for i in items), (
109+
"_process_items requires all items to share the same event_type; "
110+
"callers must split START and END batches before dispatching."
111+
)
112+
by_processor: dict[AsyncTracingProcessor, list[Span]] = {}
113+
for item in items:
114+
for p in item.processors:
115+
by_processor.setdefault(p, []).append(item.span)
116+
117+
async def _handle(p: AsyncTracingProcessor, spans: list[Span]) -> None:
101118
try:
102-
if item.event_type == SpanEventType.START:
103-
coros = [p.on_span_start(item.span) for p in item.processors]
119+
if event_type == SpanEventType.START:
120+
await p.on_spans_start(spans)
104121
else:
105-
coros = [p.on_span_end(item.span) for p in item.processors]
106-
results = await asyncio.gather(*coros, return_exceptions=True)
107-
for result in results:
108-
if isinstance(result, Exception):
109-
logger.error(
110-
"Tracing processor error during %s for span %s",
111-
item.event_type.value,
112-
item.span.id,
113-
exc_info=result,
114-
)
122+
await p.on_spans_end(spans)
115123
except Exception:
116124
logger.exception(
117-
"Unexpected error in span queue for span %s", item.span.id
125+
"Tracing processor %s failed handling %d spans during %s",
126+
type(p).__name__,
127+
len(spans),
128+
event_type.value,
118129
)
119130

120-
await asyncio.gather(*[_handle(item) for item in items])
131+
await asyncio.gather(*[_handle(p, spans) for p, spans in by_processor.items()])
121132

122133
# ------------------------------------------------------------------
123134
# Shutdown

‎tests/lib/core/tracing/processors/test_sgp_tracing_processor.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -188,3 +188,42 @@ async def test_sgp_span_input_updated_on_end(self):
188188
assert len(processor._spans) == 0
189189
# The end upsert should have been called
190190
assert processor.sgp_async_client.spans.upsert_batch.call_count == 2 # start + end
191+
192+
async def test_on_spans_start_sends_single_upsert_for_batch(self):
193+
"""Given N spans at once, on_spans_start should make ONE upsert_batch HTTP call."""
194+
processor, _ = self._make_processor()
195+
196+
n = 10
197+
spans = [_make_span() for _ in range(n)]
198+
with patch(f"{MODULE}.create_span", side_effect=lambda **kw: _make_mock_sgp_span()):
199+
await processor.on_spans_start(spans)
200+
201+
assert processor.sgp_async_client.spans.upsert_batch.call_count == 1, (
202+
"Batched on_spans_start must make exactly one upsert_batch HTTP call"
203+
)
204+
items = processor.sgp_async_client.spans.upsert_batch.call_args.kwargs["items"]
205+
assert len(items) == n
206+
# All spans should be tracked for the subsequent end call
207+
assert len(processor._spans) == n
208+
209+
async def test_on_spans_end_sends_single_upsert_for_batch(self):
210+
"""Given N spans at once, on_spans_end should make ONE upsert_batch HTTP call."""
211+
processor, _ = self._make_processor()
212+
213+
n = 10
214+
spans = [_make_span() for _ in range(n)]
215+
with patch(f"{MODULE}.create_span", side_effect=lambda **kw: _make_mock_sgp_span()):
216+
await processor.on_spans_start(spans)
217+
218+
processor.sgp_async_client.spans.upsert_batch.reset_mock()
219+
220+
for span in spans:
221+
span.end_time = datetime.now(UTC)
222+
await processor.on_spans_end(spans)
223+
224+
assert processor.sgp_async_client.spans.upsert_batch.call_count == 1, (
225+
"Batched on_spans_end must make exactly one upsert_batch HTTP call"
226+
)
227+
items = processor.sgp_async_client.spans.upsert_batch.call_args.kwargs["items"]
228+
assert len(items) == n
229+
assert len(processor._spans) == 0
Lines changed: 98 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,98 @@
1+
from __future__ import annotations
2+
3+
import uuid
4+
import logging
5+
from typing import override
6+
from datetime import UTC, datetime
7+
8+
from agentex.types.span import Span
9+
from agentex.lib.types.tracing import TracingProcessorConfig
10+
from agentex.lib.core.tracing.processors.tracing_processor_interface import (
11+
AsyncTracingProcessor,
12+
)
13+
14+
15+
def _make_span(span_id: str | None = None) -> Span:
16+
return Span(
17+
id=span_id or str(uuid.uuid4()),
18+
name="test-span",
19+
start_time=datetime.now(UTC),
20+
trace_id="trace-1",
21+
)
22+
23+
24+
class _RecordingProcessor(AsyncTracingProcessor):
25+
"""Test processor that records every on_span_* call and fails on demand."""
26+
27+
def __init__(self, fail_ids: set[str] | None = None) -> None:
28+
self.started_ids: list[str] = []
29+
self.ended_ids: list[str] = []
30+
self._fail_ids = fail_ids or set()
31+
32+
@override
33+
async def on_span_start(self, span: Span) -> None:
34+
self.started_ids.append(span.id)
35+
if span.id in self._fail_ids:
36+
raise RuntimeError(f"boom-start-{span.id}")
37+
38+
@override
39+
async def on_span_end(self, span: Span) -> None:
40+
self.ended_ids.append(span.id)
41+
if span.id in self._fail_ids:
42+
raise RuntimeError(f"boom-end-{span.id}")
43+
44+
@override
45+
async def shutdown(self) -> None:
46+
pass
47+
48+
49+
class TestDefaultBatchedFanout:
50+
"""The default on_spans_start / on_spans_end in AsyncTracingProcessor must:
51+
- dispatch to the single-span method for every span
52+
- continue after individual failures (not short-circuit)
53+
- log each failure individually
54+
- not propagate exceptions to the caller
55+
"""
56+
57+
async def test_on_spans_start_runs_every_span_despite_failures(self, caplog):
58+
proc = _RecordingProcessor(fail_ids={"span-1"})
59+
spans = [_make_span(f"span-{i}") for i in range(3)]
60+
61+
with caplog.at_level(logging.ERROR):
62+
# Must not raise, even though span-1 fails.
63+
await proc.on_spans_start(spans)
64+
65+
# Every span's on_span_start was invoked
66+
assert proc.started_ids == ["span-0", "span-1", "span-2"]
67+
68+
async def test_on_spans_start_logs_each_failure(self, caplog):
69+
proc = _RecordingProcessor(fail_ids={"span-0", "span-2"})
70+
spans = [_make_span(f"span-{i}") for i in range(3)]
71+
72+
with caplog.at_level(logging.ERROR):
73+
await proc.on_spans_start(spans)
74+
75+
# Two distinct error log records, one per failing span
76+
error_records = [r for r in caplog.records if r.levelno == logging.ERROR]
77+
messages = " ".join(r.getMessage() for r in error_records)
78+
assert "span-0" in messages
79+
assert "span-2" in messages
80+
81+
async def test_on_spans_end_runs_every_span_despite_failures(self, caplog):
82+
proc = _RecordingProcessor(fail_ids={"span-1"})
83+
spans = [_make_span(f"span-{i}") for i in range(3)]
84+
85+
with caplog.at_level(logging.ERROR):
86+
await proc.on_spans_end(spans)
87+
88+
assert proc.ended_ids == ["span-0", "span-1", "span-2"]
89+
90+
async def test_dummy_config_construction(self):
91+
"""AsyncTracingProcessor's __init__ is abstract — verify concrete
92+
subclass above satisfies the interface."""
93+
_ = TracingProcessorConfig
94+
proc = _RecordingProcessor()
95+
await proc.on_spans_start([])
96+
await proc.on_spans_end([])
97+
assert proc.started_ids == []
98+
assert proc.ended_ids == []

0 commit comments

Comments
 (0)