Skip to content

Commit 1975532

Browse files
committed
fix(traces): guard span writes and end() with the handle's lock
A same-key write from two threads could reserve two slots, and a write that lost the race to end() could land after the record was taken. The cap check, the event count and the end() snapshot now run under the handle's lock; the value walk still runs outside it.
1 parent fbf8dd7 commit 1975532

3 files changed

Lines changed: 140 additions & 47 deletions

File tree

‎posthog/test/tracing/test_span.py‎

Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import gc
44
import sys
55
import threading
6+
import time
67
import weakref
78
from contextvars import ContextVar
89
from datetime import datetime, timezone
@@ -823,6 +824,68 @@ def test_add_event_with_a_non_mapping_logs_and_keeps_the_event(self, caplog):
823824
assert "expected a mapping" in caplog.text
824825

825826

827+
class TestConcurrentWrites:
828+
def _run(self, worker, count=8):
829+
threads = [threading.Thread(target=worker, args=(i,)) for i in range(count)]
830+
for thread in threads:
831+
thread.start()
832+
for thread in threads:
833+
thread.join(30)
834+
835+
def test_writes_to_one_key_from_many_threads_spend_one_slot(self):
836+
records: list = []
837+
span = make_span(records, max_attributes=1)
838+
839+
def worker(i):
840+
for n in range(300):
841+
span.set_attribute("k", n)
842+
843+
self._run(worker)
844+
span.end()
845+
assert list(records[0].attributes) == ["k"]
846+
assert records[0].dropped_attributes_count == 0
847+
848+
def test_events_from_many_threads_never_exceed_the_cap(self):
849+
records: list = []
850+
span = make_span(records, max_events=100)
851+
852+
def worker(i):
853+
for n in range(50):
854+
span.add_event(f"{i}-{n}")
855+
856+
self._run(worker)
857+
span.end()
858+
assert len(records[0].events) == 100
859+
assert records[0].dropped_events_count == 300
860+
861+
def test_a_write_racing_end_never_lands_after_the_record(self):
862+
records: list = []
863+
span = make_span(records)
864+
stop = threading.Event()
865+
866+
def writer(i):
867+
n = 0
868+
while not stop.is_set():
869+
span.set_attribute("k", n)
870+
span.add_event("e")
871+
n += 1
872+
873+
threads = [threading.Thread(target=writer, args=(i,)) for i in range(4)]
874+
for thread in threads:
875+
thread.start()
876+
time.sleep(0.02)
877+
span.end()
878+
stop.set()
879+
for thread in threads:
880+
thread.join(30)
881+
record = records[0]
882+
assert len(records) == 1
883+
assert span._events == record.events
884+
assert {k: v for k, v in span._attributes.items() if v is not None} == (
885+
record.attributes
886+
)
887+
888+
826889
class TestExceptionStacktrace:
827890
def test_record_exception_attaches_the_stack_of_a_raised_exception(self):
828891
records: list = []

‎posthog/tracing/_otlp.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -290,7 +290,9 @@ def _to_otlp_event(event: SpanEventRecord) -> dict:
290290
attributes, cut = encode_attributes(event.attributes)
291291
if attributes:
292292
encoded["attributes"] = attributes
293-
dropped = non_negative_count(non_negative_count(event.dropped_attributes_count) + cut)
293+
dropped = non_negative_count(
294+
non_negative_count(event.dropped_attributes_count) + cut
295+
)
294296
if dropped:
295297
encoded["droppedAttributesCount"] = dropped
296298
return encoded

‎posthog/tracing/_span.py‎

Lines changed: 74 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -74,19 +74,20 @@ class _Activatable:
7474

7575
_active_var: Optional[ContextVar]
7676
_tokens: List[Token]
77-
_tokens_lock: threading.Lock
77+
# Guards the tokens and, on a recording span, every write and end().
78+
_lock: threading.Lock
7879

7980
def _activate(self) -> None:
8081
if self._active_var is not None:
81-
with self._tokens_lock:
82+
with self._lock:
8283
self._tokens.append(self._active_var.set(self))
8384

8485
def _deactivate(self) -> None:
8586
if self._active_var is None:
8687
return
8788
# The same handle can be entered in several threads or tasks at once,
8889
# and a token only resets in the context that created it.
89-
with self._tokens_lock:
90+
with self._lock:
9091
for index in range(len(self._tokens) - 1, -1, -1):
9192
try:
9293
self._active_var.reset(self._tokens[index])
@@ -117,7 +118,7 @@ def __init__(
117118
self._tracestate = tracestate
118119
self._active_var = active_var
119120
self._tokens = []
120-
self._tokens_lock = threading.Lock()
121+
self._lock = threading.Lock()
121122

122123
def traceparent(self) -> Optional[str]:
123124
return self._traceparent
@@ -244,7 +245,7 @@ def __init__(
244245
self._on_end = on_end
245246
self._active_var = active_var
246247
self._tokens = []
247-
self._tokens_lock = threading.Lock()
248+
self._lock = threading.Lock()
248249

249250
self._name = name
250251
self._kind = kind
@@ -258,11 +259,11 @@ def __init__(
258259
self._dropped_attributes = 0
259260
self._dropped_events = 0
260261
self._attributes: Dict[str, Any] = {}
261-
for key, value in (attributes or {}).items():
262-
self._write_attribute(key, value)
263262
self._events: List[SpanEventRecord] = []
264263
self._status: Optional[SpanStatus] = None
265264
self._ended = False
265+
for key, value in (attributes or {}).items():
266+
self._write_attribute(key, value)
266267

267268
def _now_ns(self) -> int:
268269
"""Now, on this span's clock basis: start plus monotonic elapsed, else wall clock."""
@@ -282,21 +283,27 @@ def _write_attribute(self, key: str, value: Any) -> None:
282283
# The encoder drops it, so it must not spend a slot.
283284
log.debug("Dropping an attribute with an empty key")
284285
return
285-
if value is None:
286-
# None removes the key, freeing its slot.
287-
if key in self._attributes and key not in self._auto_keys:
288-
self._user_attribute_count -= 1
289-
self._attributes.pop(key, None)
290-
return
291-
# Checked before the value is walked, which is the costly part.
292-
if key not in self._auto_keys and key not in self._attributes:
293-
if self._user_attribute_count >= self._max_attributes:
294-
self._dropped_attributes += 1
286+
with self._lock:
287+
if self._ended:
295288
return
296-
self._user_attribute_count += 1
297-
self._attributes[key] = truncate_attribute_value(
298-
value, self._max_attribute_value_length
299-
)
289+
if value is None:
290+
# None removes the key, freeing its slot.
291+
if key in self._attributes and key not in self._auto_keys:
292+
self._user_attribute_count -= 1
293+
self._attributes.pop(key, None)
294+
return
295+
# Checked before the value is walked, which is the costly part.
296+
if key not in self._auto_keys and key not in self._attributes:
297+
if self._user_attribute_count >= self._max_attributes:
298+
self._dropped_attributes += 1
299+
return
300+
self._user_attribute_count += 1
301+
# Reserved, so a concurrent write to the same key sees it taken.
302+
self._attributes[key] = None
303+
bounded = truncate_attribute_value(value, self._max_attribute_value_length)
304+
with self._lock:
305+
if not self._ended:
306+
self._attributes[key] = bounded
300307

301308
def set_attribute(self, key: str, value: Any) -> "Span":
302309
if self._mutable("set_attribute"):
@@ -318,11 +325,14 @@ def add_event(
318325
timestamp: Optional[SpanTimeInput] = None,
319326
) -> "Span":
320327
if self._mutable("add_event"):
321-
# A recorded exception spends a slot like any other event.
322-
if self._user_event_count >= self._max_events:
323-
self._dropped_events += 1
324-
return self
325-
self._user_event_count += 1
328+
with self._lock:
329+
if self._ended:
330+
return self
331+
# A recorded exception spends a slot like any other event.
332+
if self._user_event_count >= self._max_events:
333+
self._dropped_events += 1
334+
return self
335+
self._user_event_count += 1
326336
bounded: Optional[Dict[str, Any]] = None
327337
dropped = 0
328338
if attributes is not None:
@@ -331,18 +341,19 @@ def add_event(
331341
MAX_ATTRIBUTES_PER_EVENT,
332342
self._max_attribute_value_length,
333343
)
334-
self._events.append(
335-
SpanEventRecord(
336-
name=sanitize_name(
337-
name, "Span event name", self._max_attribute_value_length
338-
),
339-
timestamp_ns=resolve_supplied_ns(
340-
timestamp, self._now_ns(), "event timestamp"
341-
),
342-
attributes=bounded,
343-
dropped_attributes_count=dropped,
344-
)
344+
event = SpanEventRecord(
345+
name=sanitize_name(
346+
name, "Span event name", self._max_attribute_value_length
347+
),
348+
timestamp_ns=resolve_supplied_ns(
349+
timestamp, self._now_ns(), "event timestamp"
350+
),
351+
attributes=bounded,
352+
dropped_attributes_count=dropped,
345353
)
354+
with self._lock:
355+
if not self._ended:
356+
self._events.append(event)
346357
return self
347358

348359
def set_status(self, code: str, message: Optional[str] = None) -> "Span":
@@ -351,12 +362,15 @@ def set_status(self, code: str, message: Optional[str] = None) -> "Span":
351362
log.debug('Ignoring an unknown span status; expected "ok" or "error"')
352363
return self
353364
text = None if message is None else safe_str(message)
354-
self._status = SpanStatus(
365+
status = SpanStatus(
355366
code,
356367
truncate_string(text, self._max_attribute_value_length)
357368
if text
358369
else None,
359370
)
371+
with self._lock:
372+
if not self._ended:
373+
self._status = status
360374
return self
361375

362376
@property
@@ -379,9 +393,12 @@ def _record_exception(self, exception: BaseException, keep_ok: bool) -> None:
379393

380394
def update_name(self, name: str) -> "Span":
381395
if self._mutable("update_name"):
382-
self._name = sanitize_name(
396+
sanitized = sanitize_name(
383397
name, "Span name", self._max_attribute_value_length
384398
)
399+
with self._lock:
400+
if not self._ended:
401+
self._name = sanitized
385402
return self
386403

387404
def traceparent(self) -> Optional[str]:
@@ -400,11 +417,22 @@ def _child_context(self) -> ParentContext:
400417
)
401418

402419
def end(self, end_time: Optional[SpanTimeInput] = None) -> None:
403-
with self._tokens_lock:
420+
with self._lock:
404421
if self._ended:
405422
log.debug("Ignoring end() on a span that has already ended")
406423
return
407424
self._ended = True
425+
# Snapshotted under the lock: a write that lost the race to end()
426+
# must not land in the record or after it.
427+
attributes = {
428+
key: value
429+
for key, value in self._attributes.items()
430+
if value is not None
431+
}
432+
events = list(self._events)
433+
name, status = self._name, self._status
434+
dropped_attributes = self._dropped_attributes
435+
dropped_events = self._dropped_events
408436

409437
derived = self._now_ns()
410438
resolved = resolve_supplied_ns(end_time, derived, "end time")
@@ -415,15 +443,15 @@ def end(self, end_time: Optional[SpanTimeInput] = None) -> None:
415443
trace_state=self._trace_state,
416444
trace_flags=self._trace_flags,
417445
parent_is_remote=self._parent_is_remote,
418-
name=self._name,
446+
name=name,
419447
kind=self._kind,
420-
status=self._status,
421-
attributes=dict(self._attributes),
422-
events=self._events,
448+
status=status,
449+
attributes=attributes,
450+
events=events,
423451
start_ns=self._start_ns,
424452
end_ns=clamp_end_ns(resolved, self._start_ns),
425-
dropped_attributes_count=self._dropped_attributes,
426-
dropped_events_count=self._dropped_events,
453+
dropped_attributes_count=dropped_attributes,
454+
dropped_events_count=dropped_events,
427455
)
428456
try:
429457
self._on_end(record)

0 commit comments

Comments
 (0)