From 45dfcad1360fc47a8b9119651c23f131f7d55ad1 Mon Sep 17 00:00:00 2001 From: Peng Ding Date: Tue, 15 Sep 2026 01:39:57 -0500 Subject: [PATCH 1/3] feat(profiler): add TracingProfiler with per-call thread-aware tracing Add TracingProfiler class that collects per-call (function, thread_id, start_ns, end_ns) data using sys.monitoring (PEP 669) on Python 3.12+ with sys.settrace fallback for older versions. Key features: - Thread-aware collection via per-thread call stacks - Coroutine-aware: PY_YIELD/PY_RESUME produce separate spans - Same output interface as Profiler (output_text, output_html) - Raw trace access via traces() method - Aggregation to existing table/flamegraph/icicle visualizations Also extracts HTML builders to module-level functions (_render_table_html, _render_flame_html) for reuse by both Profiler and TracingProfiler. Reference library: yappi (added to bench-profiler extra). Closes #168 --- manifest.json | 12 +- profiler/profiler.py | 725 ++++++++++++++++++++++---- profiler/test_profiler_benchmark.py | 106 +++- profiler/test_profiler_correctness.py | 643 ++++++++++++++++++++++- pyproject.toml | 1 + 5 files changed, 1375 insertions(+), 112 deletions(-) diff --git a/manifest.json b/manifest.json index 4f76842..df08302 100644 --- a/manifest.json +++ b/manifest.json @@ -1,6 +1,6 @@ { "version": "1", - "generated": "2026-09-15T05:11:46.731585+00:00", + "generated": "2026-09-15T06:34:55.886016+00:00", "modules": { "a2a": { "description": "A2A (Agent-to-Agent Protocol) - Zero-dependency Python implementation", @@ -179,8 +179,8 @@ "deps": [], "tier": "subsystem", "category": "network", - "last_updated": "2026-09-14T19:04:23-05:00", - "content_hash": "4983837f8402eba1a9f7f5f5a48c7ee4eac27f893a35119c888a3d78f8c69465" + "last_updated": "2026-09-14T19:10:56-05:00", + "content_hash": "bf0527b9d00588dbe19bca5a1325c9f155f448d8c57013764911b35d232464e9" }, "jsonrpc": { "description": "JSON-RPC 2.0 -- Zero-dependency Python implementation", @@ -279,16 +279,16 @@ "content_hash": "a69b7f4a781e5bd8b8ac558c34b068800ee9e3042c28bf7c1e4a10bf672827e4" }, "profiler": { - "description": "Ergonomic cProfile wrapper with text and HTML report output", + "description": "Ergonomic profiler with text and HTML report output", "files": [ "profiler/profiler.py" ], - "version": "0.0.0", + "version": "0.1.0", "deps": [], "tier": "medium", "category": "devtools", "last_updated": "2026-09-14T20:53:33-05:00", - "content_hash": "b9bd9b27691eae5384bc5deceb1b3d008bf9f33b964bb0d28f208a8b5bdcf01a" + "content_hash": "a870e2bb6d7ff5105caeef234f21b061b75ee52e027d1ec3f029134073985a95" }, "prompt": { "description": "Zero-dependency interactive CLI prompts (confirm, select, text)", diff --git a/profiler/profiler.py b/profiler/profiler.py index be071df..da55073 100644 --- a/profiler/profiler.py +++ b/profiler/profiler.py @@ -1,12 +1,12 @@ # /// zerodep -# version = "0.0.0" +# version = "0.1.0" # deps = [] # tier = "medium" # category = "devtools" # note = "Install/update via: https://zerodep.readthedocs.io/en/latest/guide/cli/" # /// -"""Ergonomic cProfile wrapper with text and HTML report output. +"""Ergonomic profiler with text and HTML report output. Wraps ``cProfile``/``pstats`` in sync and async context managers and renders profiling data as formatted text or self-contained HTML @@ -31,6 +31,16 @@ await do_async_work() html = p.output_html() + +Thread-aware tracing:: + + from profiler import TracingProfiler + + with TracingProfiler() as p: + do_work() + + print(p.output_text()) + records = p.traces() # raw per-call data """ from __future__ import annotations @@ -39,10 +49,27 @@ import html as _html import io import pstats +import sys +import threading +import time +from collections import defaultdict from pathlib import Path -from typing import Any, cast +from typing import Any, NamedTuple, cast + +__all__ = ["Profiler", "TracingProfiler", "TraceRecord", "ProfilerError"] + + +class TraceRecord(NamedTuple): + """A single function call/return span.""" + + func: str + file: str + lineno: int + thread_id: int + start_ns: int + end_ns: int + depth: int -__all__ = ["Profiler", "ProfilerError"] _SORT_KEYS = { "cumulative": "cumulative", @@ -57,6 +84,20 @@ _VALID_STYLES = {"table", "flamegraph", "icicle"} +_SELF_FILE = __file__ + + +def _acquire_tool_id(name: str) -> int: + m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + for tid in (3, 4, 0, 1, 5): + if m.get_tool(tid) is None: + try: + m.use_tool_id(tid, name) + return tid + except ValueError: + continue + raise ProfilerError("no free sys.monitoring tool ID available") + class ProfilerError(Exception): """Raised when profiler operations fail.""" @@ -391,118 +432,594 @@ def _build_flame_html( *, inverted: bool = False, ) -> str: - """Build a self-contained flamegraph or icicle chart HTML document.""" - import json as _json - - # Use CPU total_tt (same as _extract_call_tree) for consistency total_time = cast(Any, self._stats).total_tt or 1e-9 - style_name = "Icicle Chart" if inverted else "Flamegraph" - - return ( - "\n" - '\n\n' - '\n' - '\n' - f"{_html.escape(title)}\n" - f"\n" - "\n\n" - '
\n' - f"

{_html.escape(title)}

\n" - f'

{_html.escape(style_name)} · ' - f"Total time: {total_time:.6f}s

\n" - '
\n' - '\n' - '\n' - '\n" - "
\n" - f'
\n' - "
\n
\n" - "\n" - "\n" - ) + return _render_flame_html(tree, title, total_time, inverted=inverted) def _build_table_html(self, rows: list[dict[str, Any]], title: str) -> str: - max_cumtime = max((r["cumtime"] for r in rows), default=1.0) or 1e-9 - max_tottime = max((r["tottime"] for r in rows), default=1.0) or 1e-9 + return _render_table_html(rows, title, self.total_time) - tbody_parts: list[str] = [] - for r in rows: - func_display = _html.escape(f"{r['file']}:{r['lineno']}({r['func']})") - cum_bar = r["cumtime"] / max_cumtime * 100 - tot_bar = r["tottime"] / max_tottime * 100 - - calls_str = ( - str(r["calls"]) - if r["calls"] == r["primitive_calls"] - else f"{r['calls']}/{r['primitive_calls']}" + +class TracingProfiler: + """Per-call tracing profiler with thread-aware collection. + + Collects ``(function, thread_id, start_ns, end_ns)`` data for every + Python function call, enabling per-call analysis and future waterfall + visualization. + + Uses ``sys.monitoring`` (PEP 669) on Python 3.12+ for low overhead, + falling back to ``sys.settrace`` on older versions. + + Args: + builtins: Reserved for future use (C-level tracing not yet supported). + sort_by: Default sort key for output methods. + async_mode: If True, the profiler can be used as an async + context manager. + """ + + def __init__( + self, + *, + builtins: bool = False, + sort_by: str = "cumulative", + async_mode: bool = False, + ) -> None: + self._async_mode = async_mode + self._builtins = builtins + self._default_sort = Profiler._resolve_sort_key(sort_by) + self._use_monitoring = hasattr(sys, "monitoring") + self._running = False + self._records: list[TraceRecord] = [] + self._records_lock = threading.Lock() + self._stacks: dict[int, list[tuple[str, str, int, int]]] = defaultdict(list) + self._tool_id: int | None = None + self._wall_start_ns: int = 0 + self._wall_end_ns: int = 0 + + # -- Lifecycle ----------------------------------------------------------- + + def start(self) -> None: + """Enable tracing. No-op if already running.""" + if self._running: + return + self._records.clear() + self._stacks.clear() + self._wall_start_ns = time.perf_counter_ns() + self._wall_end_ns = 0 + if self._use_monitoring: + self._start_monitoring() + else: + self._start_settrace() + self._running = True + + def stop(self) -> None: + """Disable tracing. No-op if not running.""" + if not self._running: + return + if self._use_monitoring: + self._stop_monitoring() + else: + self._stop_settrace() + self._wall_end_ns = time.perf_counter_ns() + self._running = False + + def reset(self) -> None: + """Clear all collected data.""" + if self._running: + self.stop() + self._records.clear() + self._stacks.clear() + self._wall_start_ns = 0 + self._wall_end_ns = 0 + + def __enter__(self) -> TracingProfiler: + self.start() + return self + + def __exit__(self, *args: object) -> None: + self.stop() + + async def __aenter__(self) -> TracingProfiler: + if not self._async_mode: + raise ProfilerError( + "async context manager requires TracingProfiler(async_mode=True)" + ) + self.start() + return self + + async def __aexit__(self, *args: object) -> None: + self.stop() + + # -- Properties ---------------------------------------------------------- + + @property + def is_running(self) -> bool: + """Whether the profiler is currently collecting data.""" + return self._running + + @property + def total_time(self) -> float: + """Total wall-clock time in seconds.""" + self._ensure_stopped() + end = self._wall_end_ns if not self._running else time.perf_counter_ns() + return (end - self._wall_start_ns) / 1e9 + + def traces(self) -> list[TraceRecord]: + """Return a copy of the raw trace records.""" + self._ensure_stopped() + return list(self._records) + + @property + def thread_ids(self) -> set[int]: + """Unique thread IDs observed during profiling.""" + self._ensure_stopped() + return {r.thread_id for r in self._records} + + def _ensure_stopped(self) -> None: + if self._wall_end_ns == 0 and not self._running: + raise ProfilerError("no profiling data - run the profiler first") + + # -- sys.monitoring backend ---------------------------------------------- + + def _start_monitoring(self) -> None: + m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + e = m.events + self._tool_id = _acquire_tool_id("zerodep.TracingProfiler") + m.register_callback(self._tool_id, e.PY_START, self._on_py_start) + m.register_callback(self._tool_id, e.PY_RETURN, self._on_py_exit) + m.register_callback(self._tool_id, e.PY_UNWIND, self._on_py_exit) + m.register_callback(self._tool_id, e.PY_YIELD, self._on_py_yield) + m.register_callback(self._tool_id, e.PY_RESUME, self._on_py_resume) + m.set_events( + self._tool_id, + e.PY_START | e.PY_RETURN | e.PY_UNWIND | e.PY_YIELD | e.PY_RESUME, + ) + + def _stop_monitoring(self) -> None: + m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + e = m.events + tid = self._tool_id + assert tid is not None + m.set_events(tid, e.NO_EVENTS) + for ev in (e.PY_START, e.PY_RETURN, e.PY_UNWIND, e.PY_YIELD, e.PY_RESUME): + m.register_callback(tid, ev, None) + m.free_tool_id(tid) + self._tool_id = None + + def _on_py_start(self, code: Any, offset: int) -> None: + tid = threading.get_ident() + self._stacks[tid].append( + ( + code.co_qualname, + code.co_filename, + code.co_firstlineno, + time.perf_counter_ns(), + ) + ) + + def _on_py_exit(self, code: Any, offset: int, *args: Any) -> None: + tid = threading.get_ident() + stack = self._stacks.get(tid) + if not stack: + return + func, file, lineno, start_ns = stack.pop() + end_ns = time.perf_counter_ns() + record = TraceRecord(func, file, lineno, tid, start_ns, end_ns, len(stack)) + with self._records_lock: + self._records.append(record) + + def _on_py_yield(self, code: Any, offset: int, retval: Any) -> None: + self._on_py_exit(code, offset, retval) + + def _on_py_resume(self, code: Any, offset: int) -> None: + self._on_py_start(code, offset) + + # -- sys.settrace fallback ----------------------------------------------- + + def _start_settrace(self) -> None: + sys.settrace(self._trace_func) + threading.settrace(self._trace_func) + + def _stop_settrace(self) -> None: + sys.settrace(None) + threading.settrace(None) + + def _trace_func(self, frame: Any, event: str, arg: Any) -> Any: + if event == "call": + code = frame.f_code + tid = threading.get_ident() + func = getattr(code, "co_qualname", code.co_name) + self._stacks[tid].append( + (func, code.co_filename, code.co_firstlineno, time.perf_counter_ns()) + ) + elif event == "return": + tid = threading.get_ident() + stack = self._stacks.get(tid) + if stack: + func, file, lineno, start_ns = stack.pop() + end_ns = time.perf_counter_ns() + record = TraceRecord( + func, file, lineno, tid, start_ns, end_ns, len(stack) + ) + with self._records_lock: + self._records.append(record) + return self._trace_func + + # -- Data extraction ----------------------------------------------------- + + def _extract_rows( + self, + *, + sort_by: str | None = None, + limit: int | None = None, + ) -> list[dict[str, Any]]: + self._ensure_stopped() + sort_key = ( + Profiler._resolve_sort_key(sort_by) if sort_by else self._default_sort + ) + total_wall = self.total_time or 1e-9 + records = [r for r in self._records if r.file != _SELF_FILE] + + by_thread: dict[int, list[TraceRecord]] = defaultdict(list) + for r in records: + by_thread[r.thread_id].append(r) + + child_sum: dict[int, float] = defaultdict(float) + + for _tid, recs in by_thread.items(): + recs.sort(key=lambda r: r.start_ns) + parent_stack: list[tuple[int, TraceRecord]] = [] + for i, rec in enumerate(recs): + while parent_stack and parent_stack[-1][1].end_ns <= rec.start_ns: + parent_stack.pop() + if parent_stack and parent_stack[-1][1].depth == rec.depth - 1: + parent_idx = parent_stack[-1][0] + child_sum[parent_idx] += (rec.end_ns - rec.start_ns) / 1e9 + parent_stack.append((id(rec), rec)) + + agg: dict[tuple[str, str, int], dict[str, Any]] = {} + for rec in records: + key = (rec.func, rec.file, rec.lineno) + duration = (rec.end_ns - rec.start_ns) / 1e9 + self_time = max(duration - child_sum.get(id(rec), 0.0), 0.0) + if key not in agg: + agg[key] = { + "func": rec.func, + "file": rec.file, + "lineno": rec.lineno, + "calls": 0, + "primitive_calls": 0, + "tottime": 0.0, + "cumtime": 0.0, + } + entry = agg[key] + entry["calls"] += 1 + entry["primitive_calls"] += 1 + entry["tottime"] += self_time + entry["cumtime"] += duration + + rows: list[dict[str, Any]] = [] + for entry in agg.values(): + nc = entry["calls"] + cc = entry["primitive_calls"] + tt = entry["tottime"] + ct = entry["cumtime"] + rows.append( + { + **entry, + "percall_tot": tt / nc if nc else 0.0, + "percall_cum": ct / cc if cc else 0.0, + "tottime_pct": tt / total_wall * 100, + "cumtime_pct": ct / total_wall * 100, + } ) - tbody_parts.append( - f"" - f'{func_display}' - f'' - f'
' - f'{r["cumtime"]:.6f}' - f'' - f'
' - f'{r["tottime"]:.6f}' - f'{calls_str}' - f'' - f"{r['percall_cum']:.6f}" - f'' - f"{r['cumtime_pct']:.1f}%" - f"" + sort_map: dict[str, Any] = { + "cumulative": lambda r: r["cumtime"], + "tottime": lambda r: r["tottime"], + "calls": lambda r: r["calls"], + "name": lambda r: r["func"], + "filename": lambda r: r["file"], + } + key_fn = sort_map.get(sort_key, sort_map["cumulative"]) + reverse = sort_key not in ("name", "filename") + rows.sort(key=key_fn, reverse=reverse) + + if limit is not None: + rows = rows[:limit] + return rows + + def _extract_call_tree(self) -> list[dict[str, Any]]: + self._ensure_stopped() + total_wall = self.total_time or 1e-9 + records = [r for r in self._records if r.file != _SELF_FILE] + + by_thread: dict[int, list[TraceRecord]] = defaultdict(list) + for r in records: + by_thread[r.thread_id].append(r) + + all_roots: list[dict[str, Any]] = [] + + for _tid, recs in by_thread.items(): + recs.sort(key=lambda r: r.start_ns) + stack: list[tuple[int, dict[str, Any]]] = [] + roots: list[dict[str, Any]] = [] + + for rec in recs: + duration = (rec.end_ns - rec.start_ns) / 1e9 + node: dict[str, Any] = { + "name": rec.func, + "file": rec.file, + "lineno": rec.lineno, + "cumtime": duration, + "tottime": 0.0, + "calls": 1, + "cumtime_pct": duration / total_wall * 100, + "children": [], + } + + while stack and stack[-1][0] >= rec.depth: + stack.pop() + + if stack: + stack[-1][1]["children"].append(node) + else: + roots.append(node) + + stack.append((rec.depth, node)) + + all_roots.extend(roots) + + def _compute_self_time(node: dict[str, Any]) -> None: + child_cum = sum(c["cumtime"] for c in node["children"]) + node["tottime"] = max(node["cumtime"] - child_cum, 0.0) + for child in node["children"]: + _compute_self_time(child) + + def _merge_children(node: dict[str, Any]) -> None: + merged: dict[str, dict[str, Any]] = {} + for child in node["children"]: + key = child["name"] + if key in merged: + merged[key]["cumtime"] += child["cumtime"] + merged[key]["tottime"] += child["tottime"] + merged[key]["calls"] += child["calls"] + merged[key]["cumtime_pct"] += child["cumtime_pct"] + merged[key]["children"].extend(child["children"]) + else: + merged[key] = child + node["children"] = sorted( + merged.values(), key=lambda c: c["cumtime"], reverse=True ) + for child in node["children"]: + _merge_children(child) + + for root in all_roots: + _compute_self_time(root) + _merge_children(root) + + all_roots.sort(key=lambda n: n["cumtime"], reverse=True) + return all_roots + + # -- Output methods ------------------------------------------------------ - total_funcs = len(rows) + def output_text( + self, + *, + sort_by: str | None = None, + limit: int | None = None, + file: str | Path | None = None, + ) -> str: + """Return formatted text profiling output. + + Args: + sort_by: Override the default sort key. + limit: Show only the top N functions. + file: If provided, also write the text to this file path. + + Returns: + Formatted text string. + """ + self._ensure_stopped() + rows = self._extract_rows(sort_by=sort_by, limit=limit) + + total_calls = sum(r["calls"] for r in rows) total_time_s = self.total_time + sort_key = ( + Profiler._resolve_sort_key(sort_by) if sort_by else self._default_sort + ) + + lines: list[str] = [ + f" {total_calls} function calls in {total_time_s:.3f} seconds\n", + f" Ordered by: {sort_key} time\n", + "", + f"{'ncalls':>9s} {'tottime':>9s} {'percall':>9s} " + f"{'cumtime':>9s} {'percall':>9s} filename:lineno(function)", + ] + for r in rows: + calls_str = str(r["calls"]) + if r["calls"] != r["primitive_calls"]: + calls_str = f"{r['calls']}/{r['primitive_calls']}" + lines.append( + f"{calls_str:>9s} {r['tottime']:>9.3f} " + f"{r['percall_tot']:>9.3f} {r['cumtime']:>9.3f} " + f"{r['percall_cum']:>9.3f} " + f"{r['file']}:{r['lineno']}({r['func']})" + ) + + text = "\n".join(lines) + "\n" + if file is not None: + Path(file).write_text(text, encoding="utf-8") + return text + + def output_html( + self, + *, + style: str = "table", + sort_by: str | None = None, + limit: int | None = None, + title: str = "Profile Report", + file: str | Path | None = None, + ) -> str: + """Return a self-contained HTML profiling report. + + Args: + style: Output style — ``"table"``, ``"flamegraph"``, or ``"icicle"``. + sort_by: Override the default sort key for initial table order. + limit: Show only the top N functions. + title: HTML page title. + file: If provided, also write the HTML to this file path. + + Returns: + Complete HTML document string with inline CSS/JS. + """ + self._ensure_stopped() + if style not in _VALID_STYLES: + raise ValueError( + f"unknown style {style!r}, expected one of: " + f"{', '.join(sorted(_VALID_STYLES))}" + ) + + if style in ("flamegraph", "icicle"): + if sort_by is not None or limit is not None: + raise ValueError("sort_by and limit only apply to style='table'") + tree = self._extract_call_tree() + doc = _render_flame_html( + tree, title, self.total_time, inverted=(style == "icicle") + ) + else: + rows = self._extract_rows(sort_by=sort_by, limit=limit) + doc = _render_table_html(rows, title, self.total_time) - return ( - "\n" - '\n\n' - '\n' - '\n' - f"{_html.escape(title)}\n" - f"\n" - "\n\n" - '
\n' - f"

{_html.escape(title)}

\n" - f'

Total time: {total_time_s:.6f}s · ' - f"{total_funcs} functions

\n" - '
\n' - '\n' - '\n" - "
\n" - '
\n' - '\n\n' - '\n' - '\n' - '\n' - '\n' - '\n' - '\n' - "\n\n" - + "\n".join(tbody_parts) - + "\n
Function' - 'Cumulative' - 'Total (self)' - 'Calls' - 'Per Call (cum)' - '% of Total' - '
\n
\n
\n" - f"\n" - "\n" + if file is not None: + Path(file).write_text(doc, encoding="utf-8") + return doc + + +# --------------------------------------------------------------------------- +# Module-level HTML rendering functions (shared by Profiler & TracingProfiler) +# --------------------------------------------------------------------------- + + +def _render_flame_html( + tree: list[dict[str, Any]], + title: str, + total_time: float, + *, + inverted: bool = False, +) -> str: + import json as _json + + total_time = total_time or 1e-9 + style_name = "Icicle Chart" if inverted else "Flamegraph" + + return ( + "\n" + '\n\n' + '\n' + '\n' + f"{_html.escape(title)}\n" + f"\n" + "\n\n" + '
\n' + f"

{_html.escape(title)}

\n" + f'

{_html.escape(style_name)} · ' + f"Total time: {total_time:.6f}s

\n" + '
\n' + '\n' + '\n' + '\n" + "
\n" + f'
\n' + "
\n
\n" + "\n" + "\n" + ) + + +def _render_table_html( + rows: list[dict[str, Any]], title: str, total_time: float +) -> str: + max_cumtime = max((r["cumtime"] for r in rows), default=1.0) or 1e-9 + max_tottime = max((r["tottime"] for r in rows), default=1.0) or 1e-9 + + tbody_parts: list[str] = [] + for r in rows: + func_display = _html.escape(f"{r['file']}:{r['lineno']}({r['func']})") + cum_bar = r["cumtime"] / max_cumtime * 100 + tot_bar = r["tottime"] / max_tottime * 100 + + calls_str = ( + str(r["calls"]) + if r["calls"] == r["primitive_calls"] + else f"{r['calls']}/{r['primitive_calls']}" ) + tbody_parts.append( + f"" + f'{func_display}' + f'' + f'
' + f'{r["cumtime"]:.6f}' + f'' + f'
' + f'{r["tottime"]:.6f}' + f'{calls_str}' + f'' + f"{r['percall_cum']:.6f}" + f'' + f"{r['cumtime_pct']:.1f}%" + f"" + ) + + total_funcs = len(rows) + total_time_s = total_time + + return ( + "\n" + '\n\n' + '\n' + '\n' + f"{_html.escape(title)}\n" + f"\n" + "\n\n" + '
\n' + f"

{_html.escape(title)}

\n" + f'

Total time: {total_time_s:.6f}s · ' + f"{total_funcs} functions

\n" + '
\n' + '\n' + '\n" + "
\n" + '
\n' + '\n\n' + '\n' + '\n' + '\n' + '\n' + '\n' + '\n' + "\n\n" + + "\n".join(tbody_parts) + + "\n
Function' + 'Cumulative' + 'Total (self)' + 'Calls' + 'Per Call (cum)' + '% of Total' + '
\n
\n
\n" + f"\n" + "\n" + ) + # --------------------------------------------------------------------------- # Inline CSS for HTML table output diff --git a/profiler/test_profiler_benchmark.py b/profiler/test_profiler_benchmark.py index 7e7afa5..f2ae2cd 100644 --- a/profiler/test_profiler_benchmark.py +++ b/profiler/test_profiler_benchmark.py @@ -8,7 +8,7 @@ sys.path.insert(0, os.path.dirname(__file__)) -from profiler import Profiler # noqa: E402 +from profiler import Profiler, TracingProfiler # noqa: E402 pyinstrument = pytest.importorskip("pyinstrument", reason="pyinstrument not installed") @@ -87,3 +87,107 @@ def test_pyinstrument_text(self, benchmark): def test_pyinstrument_html(self, benchmark): p = self._make_pyinstrument_profiler() benchmark(p.output_html) + + +# -- yappi conditional import ------------------------------------------------ + +try: + import yappi as _yappi + + _HAS_YAPPI = True +except ImportError: + _HAS_YAPPI = False + + +# -- TestTracingProfilerOverhead --------------------------------------------- + + +class TestTracingProfilerOverhead: + """Compare TracingProfiler vs yappi overhead on the same workload.""" + + N = 2000 + + def test_zerodep_tracing(self, benchmark): + def run(): + with TracingProfiler() as _p: + _workload(self.N) + + benchmark(run) + + @pytest.mark.skipif(not _HAS_YAPPI, reason="yappi not installed") + def test_yappi(self, benchmark): + def run(): + _yappi.set_clock_type("wall") + _yappi.start(builtins=False) + _workload(self.N) + _yappi.stop() + _yappi.clear_stats() + + benchmark(run) + + +# -- TestTracingOutputGeneration --------------------------------------------- + + +class TestTracingOutputGeneration: + """Compare tracing profiler output generation speed.""" + + def _make_tracing_profiler(self): + p = TracingProfiler() + p.start() + _workload(5000) + p.stop() + return p + + def test_zerodep_tracing_text(self, benchmark): + p = self._make_tracing_profiler() + benchmark(p.output_text) + + def test_zerodep_tracing_html_table(self, benchmark): + p = self._make_tracing_profiler() + benchmark(p.output_html, style="table") + + def test_zerodep_tracing_html_flamegraph(self, benchmark): + p = self._make_tracing_profiler() + benchmark(p.output_html, style="flamegraph") + + +# -- TestMultiThreadOverhead ------------------------------------------------- + +import threading + + +class TestMultiThreadOverhead: + """Compare multi-threaded profiling overhead.""" + + N = 1000 + + def test_zerodep_tracing_multithread(self, benchmark): + def run(): + with TracingProfiler() as _p: + threads = [] + for _ in range(4): + t = threading.Thread(target=_workload, args=(self.N,)) + t.start() + threads.append(t) + for t in threads: + t.join() + + benchmark(run) + + @pytest.mark.skipif(not _HAS_YAPPI, reason="yappi not installed") + def test_yappi_multithread(self, benchmark): + def run(): + _yappi.set_clock_type("wall") + _yappi.start(builtins=False) + threads = [] + for _ in range(4): + t = threading.Thread(target=_workload, args=(self.N,)) + t.start() + threads.append(t) + for t in threads: + t.join() + _yappi.stop() + _yappi.clear_stats() + + benchmark(run) diff --git a/profiler/test_profiler_correctness.py b/profiler/test_profiler_correctness.py index 651b8a0..b128e70 100644 --- a/profiler/test_profiler_correctness.py +++ b/profiler/test_profiler_correctness.py @@ -4,12 +4,13 @@ import os import sys import tempfile +import threading import pytest sys.path.insert(0, os.path.dirname(__file__)) -from profiler import Profiler, ProfilerError # noqa: E402 +from profiler import Profiler, ProfilerError, TraceRecord, TracingProfiler # noqa: E402 # -- Helpers ----------------------------------------------------------------- @@ -32,6 +33,14 @@ async def _async_work() -> int: return _busy_work(500) +try: + import yappi as _yappi + + _HAS_YAPPI = True +except ImportError: + _HAS_YAPPI = False + + # -- TestSyncProfiler ------------------------------------------------------- @@ -640,3 +649,635 @@ def test_icicle_to_file(self): assert os.path.isfile(path) with open(path, encoding="utf-8") as f: assert f.read() == html + + +# -- Yappi availability flag ------------------------------------------------- + +try: + import yappi as _yappi + + _HAS_YAPPI = True +except ImportError: + _HAS_YAPPI = False + + +# ============================================================================ +# TracingProfiler tests +# ============================================================================ + + +# -- TestTracingSyncProfiler ------------------------------------------------- + + +class TestTracingSyncProfiler: + """Sync context manager and start/stop lifecycle for TracingProfiler.""" + + def test_context_manager_basic(self): + with TracingProfiler() as p: + _busy_work() + assert not p.is_running + + def test_start_stop_explicit(self): + p = TracingProfiler() + p.start() + assert p.is_running + _busy_work() + p.stop() + assert not p.is_running + + def test_double_start_is_noop(self): + p = TracingProfiler() + p.start() + p.start() # should not corrupt state + _busy_work() + p.stop() + assert not p.is_running + + def test_double_stop_is_noop(self): + p = TracingProfiler() + p.start() + _busy_work() + p.stop() + p.stop() # should be safe + assert not p.is_running + + def test_stop_without_start_is_safe(self): + p = TracingProfiler() + p.stop() # no crash + + def test_is_running_property(self): + p = TracingProfiler() + assert not p.is_running + p.start() + assert p.is_running + p.stop() + assert not p.is_running + + def test_reset_clears_data(self): + with TracingProfiler() as p: + _busy_work() + p.reset() + with pytest.raises(ProfilerError): + p.traces() + + def test_reuse_after_reset(self): + p = TracingProfiler() + with p: + _busy_work() + p.reset() + with p: + _busy_work() + records = p.traces() + assert len(records) > 0 + + def test_context_manager_returns_self(self): + tp = TracingProfiler() + with tp as p: + _busy_work() + assert p is tp + + +# -- TestTracingAsyncProfiler ------------------------------------------------ + + +class TestTracingAsyncProfiler: + """Async context manager for TracingProfiler.""" + + @pytest.mark.asyncio + async def test_async_context_manager(self): + async with TracingProfiler(async_mode=True) as p: + await _async_work() + assert not p.is_running + + @pytest.mark.asyncio + async def test_async_mode_required(self): + with pytest.raises(ProfilerError, match="async_mode=True"): + async with TracingProfiler(): + await _async_work() + + @pytest.mark.asyncio + async def test_sync_with_async_mode_true_still_works(self): + with TracingProfiler(async_mode=True) as p: + _busy_work() + assert not p.is_running + + +# -- TestTraceData ----------------------------------------------------------- + + +class TestTraceData: + """TracingProfiler-specific trace record tests.""" + + def test_traces_returns_list_of_records(self): + with TracingProfiler() as p: + _busy_work() + records = p.traces() + assert isinstance(records, list) + for r in records: + assert isinstance(r, TraceRecord) + + def test_trace_record_fields(self): + with TracingProfiler() as p: + _busy_work() + records = p.traces() + assert len(records) > 0 + for r in records: + assert isinstance(r.func, str) + assert isinstance(r.file, str) + assert isinstance(r.lineno, int) + assert isinstance(r.thread_id, int) + assert isinstance(r.start_ns, int) + assert isinstance(r.end_ns, int) + assert isinstance(r.depth, int) + + def test_start_before_end(self): + with TracingProfiler() as p: + _busy_work() + for r in p.traces(): + assert r.start_ns < r.end_ns + + def test_depth_zero_for_top_level(self): + with TracingProfiler() as p: + _busy_work() + records = p.traces() + assert any(r.depth == 0 for r in records) + + def test_depth_increases_for_nested(self): + with TracingProfiler() as p: + _call_tree_workload() + records = p.traces() + assert any(r.depth > 0 for r in records) + + def test_profiled_function_appears(self): + with TracingProfiler() as p: + _busy_work() + names = [r.func for r in p.traces()] + assert any("_busy_work" in n for n in names) + + def test_call_count_matches(self): + with TracingProfiler() as p: + for _ in range(5): + _busy_work(10) + records = p.traces() + count = sum(1 for r in records if "_busy_work" in r.func) + assert count == 5 + + def test_internal_frames_excluded(self): + profiler_path = os.path.join(os.path.dirname(__file__), "profiler.py") + with TracingProfiler() as p: + _busy_work() + for r in p.traces(): + assert r.file != profiler_path + + def test_recursive_function_traces(self): + with TracingProfiler() as p: + _recursive_fib(5) + records = p.traces() + fib_count = sum(1 for r in records if "_recursive_fib" in r.func) + assert fib_count == 15 + + +# -- TestThreadAwareness ----------------------------------------------------- + + +class TestThreadAwareness: + """Multi-thread tracing tests for TracingProfiler.""" + + def test_multi_thread_traces(self): + with TracingProfiler() as p: + _busy_work(100) # main thread work + threads = [ + threading.Thread(target=_busy_work, args=(100,)) for _ in range(2) + ] + for t in threads: + t.start() + for t in threads: + t.join() + assert len(p.thread_ids) >= 2 + + def test_per_thread_records(self): + with TracingProfiler() as p: + _busy_work(100) + threads = [ + threading.Thread(target=_busy_work, args=(100,)) for _ in range(2) + ] + for t in threads: + t.start() + for t in threads: + t.join() + records = p.traces() + by_tid: dict[int, list[TraceRecord]] = {} + for r in records: + by_tid.setdefault(r.thread_id, []).append(r) + for tid, recs in by_tid.items(): + assert all(r.thread_id == tid for r in recs) + + def test_thread_id_matches_threading_get_ident(self): + main_tid = threading.get_ident() + with TracingProfiler() as p: + _busy_work(100) + records = p.traces() + busy_recs = [r for r in records if "_busy_work" in r.func] + assert any(r.thread_id == main_tid for r in busy_recs) + + +# -- TestCoroutineTracing ---------------------------------------------------- + + +class TestCoroutineTracing: + """Async/coroutine tracing tests for TracingProfiler.""" + + @pytest.mark.asyncio + async def test_async_captures_coroutine(self): + async with TracingProfiler(async_mode=True) as p: + await _async_work() + records = p.traces() + names = [r.func for r in records] + assert any("_async_work" in n for n in names) + + @pytest.mark.asyncio + async def test_coroutine_multiple_spans(self): + async def _yielding(): + await asyncio.sleep(0) + return 1 + + async with TracingProfiler(async_mode=True) as p: + await _yielding() + records = p.traces() + yield_recs = [r for r in records if "_yielding" in r.func] + assert len(yield_recs) >= 1 + + +# -- TestTracingTextOutput --------------------------------------------------- + + +class TestTracingTextOutput: + """output_text() tests for TracingProfiler.""" + + def test_output_text_basic(self): + with TracingProfiler() as p: + _busy_work() + text = p.output_text() + assert isinstance(text, str) + assert len(text) > 0 + assert "function calls" in text + + def test_output_text_contains_function(self): + with TracingProfiler() as p: + _busy_work() + text = p.output_text() + assert "_busy_work" in text + + def test_output_text_sort_by(self): + with TracingProfiler() as p: + _busy_work() + text = p.output_text(sort_by="tottime") + assert len(text) > 0 + + def test_output_text_limit(self): + with TracingProfiler() as p: + _call_tree_workload() + text = p.output_text(limit=3) + # Count data lines (skip header lines) + lines = text.strip().split("\n") + data_lines = [ + ln + for ln in lines + if ln.strip() + and not ln.strip().startswith(("Ordered", "ncalls")) + and "function calls" not in ln + ] + assert len(data_lines) <= 3 + + def test_output_text_to_file(self): + with TracingProfiler() as p: + _busy_work() + with tempfile.TemporaryDirectory() as d: + path = os.path.join(d, "out.txt") + text = p.output_text(file=path) + assert os.path.isfile(path) + with open(path, encoding="utf-8") as f: + assert f.read() == text + + def test_output_text_before_profiling_raises(self): + p = TracingProfiler() + with pytest.raises(ProfilerError): + p.output_text() + + def test_output_text_invalid_sort_raises(self): + with TracingProfiler() as p: + _busy_work() + with pytest.raises(ValueError, match="unknown sort key"): + p.output_text(sort_by="nonexistent") + + +# -- TestTracingHtmlOutput --------------------------------------------------- + + +class TestTracingHtmlOutput: + """output_html() tests for TracingProfiler.""" + + def test_output_html_basic(self): + with TracingProfiler() as p: + _busy_work() + html = p.output_html() + assert "" in html + + def test_output_html_self_contained(self): + with TracingProfiler() as p: + _busy_work() + html = p.output_html() + assert "http://" not in html + assert "https://" not in html + + def test_output_html_table_columns(self): + with TracingProfiler() as p: + _busy_work() + html = p.output_html() + assert "Cumulative" in html + assert "Total (self)" in html + assert "Calls" in html + + def test_output_html_custom_title(self): + with TracingProfiler() as p: + _busy_work() + html = p.output_html(title="Custom Title") + assert "Custom Title" in html + + def test_output_html_to_file(self): + with TracingProfiler() as p: + _busy_work() + with tempfile.TemporaryDirectory() as d: + path = os.path.join(d, "report.html") + html = p.output_html(file=path) + assert os.path.isfile(path) + with open(path, encoding="utf-8") as f: + assert f.read() == html + + def test_output_html_unknown_style_raises(self): + with TracingProfiler() as p: + _busy_work() + with pytest.raises(ValueError, match="unknown style"): + p.output_html(style="unknown") + + def test_output_html_before_profiling_raises(self): + p = TracingProfiler() + with pytest.raises(ProfilerError): + p.output_html() + + +# -- TestTracingFlamegraphOutput --------------------------------------------- + + +class TestTracingFlamegraphOutput: + """Flamegraph output tests for TracingProfiler.""" + + def test_flamegraph_basic(self): + with TracingProfiler() as p: + _call_tree_workload() + html = p.output_html(style="flamegraph") + assert "FLAME_DATA" in html + assert "" in html + + def test_flamegraph_self_contained(self): + with TracingProfiler() as p: + _call_tree_workload() + html = p.output_html(style="flamegraph") + assert "http://" not in html + assert "https://" not in html + + def test_flamegraph_sort_by_raises(self): + with TracingProfiler() as p: + _call_tree_workload() + with pytest.raises(ValueError, match="sort_by and limit"): + p.output_html(style="flamegraph", sort_by="tottime") + + +# -- TestTracingIcicleOutput ------------------------------------------------- + + +class TestTracingIcicleOutput: + """Icicle chart output tests for TracingProfiler.""" + + def test_icicle_basic(self): + with TracingProfiler() as p: + _call_tree_workload() + html = p.output_html(style="icicle") + assert "FLAME_DATA" in html + assert "" in html + + def test_icicle_is_inverted(self): + with TracingProfiler() as p: + _call_tree_workload() + html = p.output_html(style="icicle") + assert "inverted" in html + + +# -- TestTracingDataExtraction ----------------------------------------------- + + +class TestTracingDataExtraction: + """Internal _extract_rows() tests for TracingProfiler.""" + + def test_extract_rows_returns_dicts(self): + with TracingProfiler() as p: + _busy_work() + rows = p._extract_rows() + assert isinstance(rows, list) + assert len(rows) > 0 + expected_keys = { + "func", + "file", + "lineno", + "calls", + "primitive_calls", + "tottime", + "cumtime", + "percall_tot", + "percall_cum", + "tottime_pct", + "cumtime_pct", + } + for row in rows: + assert isinstance(row, dict) + assert expected_keys.issubset(row.keys()) + + def test_extract_rows_limit(self): + with TracingProfiler() as p: + _call_tree_workload() + rows = p._extract_rows(limit=2) + assert len(rows) <= 2 + + def test_profiled_function_in_rows(self): + with TracingProfiler() as p: + _busy_work() + rows = p._extract_rows() + funcs = [r["func"] for r in rows] + assert any("_busy_work" in f for f in funcs) + + def test_call_counts_in_rows(self): + with TracingProfiler() as p: + for _ in range(5): + _busy_work(10) + rows = p._extract_rows() + for row in rows: + if "_busy_work" in row["func"]: + assert row["calls"] == 5 + break + else: + pytest.fail("_busy_work not found in extracted rows") + + +# -- TestTracingCallTreeExtraction ------------------------------------------- + + +class TestTracingCallTreeExtraction: + """Internal _extract_call_tree() tests for TracingProfiler.""" + + def test_returns_list_of_root_nodes(self): + with TracingProfiler() as p: + _call_tree_workload() + tree = p._extract_call_tree() + assert isinstance(tree, list) + assert len(tree) >= 1 + + def test_root_node_has_expected_keys(self): + with TracingProfiler() as p: + _call_tree_workload() + tree = p._extract_call_tree() + expected_keys = { + "name", + "file", + "lineno", + "cumtime", + "tottime", + "calls", + "cumtime_pct", + "children", + } + for node in tree: + assert expected_keys.issubset(node.keys()) + + def test_children_are_nested(self): + with TracingProfiler() as p: + _call_tree_workload() + tree = p._extract_call_tree() + has_children = any(len(node["children"]) > 0 for node in tree) + assert has_children + + +# -- TestTracingEdgeCases ---------------------------------------------------- + + +class TestTracingEdgeCases: + """Edge case tests for TracingProfiler.""" + + def test_empty_profile(self): + with TracingProfiler() as p: + pass + text = p.output_text() + assert isinstance(text, str) + + def test_exception_in_profiled_code(self): + p = TracingProfiler() + try: + with p: + _busy_work() + raise RuntimeError("boom") + except RuntimeError: + pass + assert not p.is_running + records = p.traces() + assert isinstance(records, list) + text = p.output_text() + assert isinstance(text, str) + + def test_very_short_execution(self): + with TracingProfiler() as p: + x = 1 # noqa: F841 + text = p.output_text() + assert isinstance(text, str) + + +# -- TestTracingVsYappi ------------------------------------------------------ + + +@pytest.mark.skipif(not _HAS_YAPPI, reason="yappi not installed") +class TestTracingVsYappi: + """Cross-validate TracingProfiler against yappi.""" + + def test_call_counts_match(self): + # TracingProfiler + with TracingProfiler() as p: + for _ in range(5): + _busy_work(100) + records = p.traces() + tp_calls = sum(1 for r in records if "_busy_work" in r.func) + + # yappi + _yappi.set_clock_type("wall") + _yappi.start(builtins=False) + for _ in range(5): + _busy_work(100) + _yappi.stop() + stats = _yappi.get_func_stats() + yappi_calls = None + for s in stats: + if "_busy_work" in s.name: + yappi_calls = s.ncall + _yappi.clear_stats() + + assert tp_calls == 5 + assert yappi_calls == 5 + + def test_function_names_match(self): + # TracingProfiler + with TracingProfiler() as p: + _busy_work(100) + tp_names = {r.func for r in p.traces()} + + # yappi + _yappi.set_clock_type("wall") + _yappi.start(builtins=False) + _busy_work(100) + _yappi.stop() + stats = _yappi.get_func_stats() + yappi_names = {s.name for s in stats} + _yappi.clear_stats() + + assert any("_busy_work" in n for n in tp_names) + assert any("_busy_work" in n for n in yappi_names) + + def test_multi_thread_call_counts(self): + # TracingProfiler + with TracingProfiler() as p: + _busy_work(100) + threads = [ + threading.Thread(target=_busy_work, args=(100,)) for _ in range(2) + ] + for t in threads: + t.start() + for t in threads: + t.join() + records = p.traces() + tp_calls = sum(1 for r in records if "_busy_work" in r.func) + + # yappi + _yappi.set_clock_type("wall") + _yappi.start(builtins=False) + _busy_work(100) + threads = [threading.Thread(target=_busy_work, args=(100,)) for _ in range(2)] + for t in threads: + t.start() + for t in threads: + t.join() + _yappi.stop() + stats = _yappi.get_func_stats() + yappi_calls = None + for s in stats: + if "_busy_work" in s.name: + yappi_calls = s.ncall + _yappi.clear_stats() + + assert tp_calls == 3 + assert yappi_calls == 3 diff --git a/pyproject.toml b/pyproject.toml index 7d088b2..742e65e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -166,6 +166,7 @@ bench-jsonschema = [ ] bench-profiler = [ "pyinstrument>=5.0.0", + "yappi>=1.6.0", ] dev = [ "ruff==0.15.20", From 200793eaf4f0e92cbb5e12aa062c8d2e4c5209ac Mon Sep 17 00:00:00 2001 From: Peng Ding Date: Tue, 15 Sep 2026 01:55:02 -0500 Subject: [PATCH 2/3] fix: address review feedback on TracingProfiler Actionable fixes: - Raise NotImplementedError when builtins=True (C-level tracing unsupported) - Include PROFILER_ID (2) in tool ID acquisition list - Extract _resolve_sort_key to module-level function - Key _merge_children by (name, file, lineno) not just name - Verify popped frame matches code in _on_py_exit - Use stable integer indices instead of id(rec) in _extract_rows - Remove duplicate yappi import block in tests Non-blocking fixes: - Guard against tool ID leak on partial _start_monitoring failure - Save/restore pre-existing trace function in settrace backend - Raise ProfilerError when accessing data while profiler is running --- profiler/profiler.py | 103 ++++++++++++++++---------- profiler/test_profiler_correctness.py | 18 ++--- 2 files changed, 72 insertions(+), 49 deletions(-) diff --git a/profiler/profiler.py b/profiler/profiler.py index da55073..473656e 100644 --- a/profiler/profiler.py +++ b/profiler/profiler.py @@ -87,9 +87,30 @@ class TraceRecord(NamedTuple): _SELF_FILE = __file__ +def _resolve_sort_key(key: str) -> str: + """Resolve a sort key alias to its canonical pstats name. + + Args: + key: Sort key string (e.g. ``"cumulative"``, ``"tottime"``). + + Returns: + Canonical sort key name. + + Raises: + ValueError: If *key* is not recognized. + """ + resolved = _SORT_KEYS.get(key) + if resolved is None: + raise ValueError( + f"unknown sort key {key!r}, expected one of: " + f"{', '.join(sorted(_SORT_KEYS))}" + ) + return resolved + + def _acquire_tool_id(name: str) -> int: m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] - for tid in (3, 4, 0, 1, 5): + for tid in (2, 3, 4, 0, 1, 5): if m.get_tool(tid) is None: try: m.use_tool_id(tid, name) @@ -142,13 +163,7 @@ def _make_profile(self) -> cProfile.Profile: @staticmethod def _resolve_sort_key(key: str) -> str: - resolved = _SORT_KEYS.get(key) - if resolved is None: - raise ValueError( - f"unknown sort key {key!r}, expected one of: " - f"{', '.join(sorted(_SORT_KEYS))}" - ) - return resolved + return _resolve_sort_key(key) # -- Lifecycle ----------------------------------------------------------- @@ -465,7 +480,9 @@ def __init__( ) -> None: self._async_mode = async_mode self._builtins = builtins - self._default_sort = Profiler._resolve_sort_key(sort_by) + if builtins: + raise NotImplementedError("C-level function tracing is not yet supported") + self._default_sort = _resolve_sort_key(sort_by) self._use_monitoring = hasattr(sys, "monitoring") self._running = False self._records: list[TraceRecord] = [] @@ -474,6 +491,7 @@ def __init__( self._tool_id: int | None = None self._wall_start_ns: int = 0 self._wall_end_ns: int = 0 + self._prev_trace: Any = None # -- Lifecycle ----------------------------------------------------------- @@ -555,7 +573,9 @@ def thread_ids(self) -> set[int]: return {r.thread_id for r in self._records} def _ensure_stopped(self) -> None: - if self._wall_end_ns == 0 and not self._running: + if self._running: + raise ProfilerError("profiler is still running - call stop() first") + if self._wall_end_ns == 0: raise ProfilerError("no profiling data - run the profiler first") # -- sys.monitoring backend ---------------------------------------------- @@ -564,15 +584,20 @@ def _start_monitoring(self) -> None: m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] e = m.events self._tool_id = _acquire_tool_id("zerodep.TracingProfiler") - m.register_callback(self._tool_id, e.PY_START, self._on_py_start) - m.register_callback(self._tool_id, e.PY_RETURN, self._on_py_exit) - m.register_callback(self._tool_id, e.PY_UNWIND, self._on_py_exit) - m.register_callback(self._tool_id, e.PY_YIELD, self._on_py_yield) - m.register_callback(self._tool_id, e.PY_RESUME, self._on_py_resume) - m.set_events( - self._tool_id, - e.PY_START | e.PY_RETURN | e.PY_UNWIND | e.PY_YIELD | e.PY_RESUME, - ) + try: + m.register_callback(self._tool_id, e.PY_START, self._on_py_start) + m.register_callback(self._tool_id, e.PY_RETURN, self._on_py_exit) + m.register_callback(self._tool_id, e.PY_UNWIND, self._on_py_exit) + m.register_callback(self._tool_id, e.PY_YIELD, self._on_py_yield) + m.register_callback(self._tool_id, e.PY_RESUME, self._on_py_resume) + m.set_events( + self._tool_id, + e.PY_START | e.PY_RETURN | e.PY_UNWIND | e.PY_YIELD | e.PY_RESUME, + ) + except BaseException: + m.free_tool_id(self._tool_id) + self._tool_id = None + raise def _stop_monitoring(self) -> None: m = sys.monitoring # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] @@ -601,7 +626,10 @@ def _on_py_exit(self, code: Any, offset: int, *args: Any) -> None: stack = self._stacks.get(tid) if not stack: return - func, file, lineno, start_ns = stack.pop() + func, file, lineno, start_ns = stack[-1] + if func != code.co_qualname: + return # stack desync — skip rather than corrupt + stack.pop() end_ns = time.perf_counter_ns() record = TraceRecord(func, file, lineno, tid, start_ns, end_ns, len(stack)) with self._records_lock: @@ -616,12 +644,13 @@ def _on_py_resume(self, code: Any, offset: int) -> None: # -- sys.settrace fallback ----------------------------------------------- def _start_settrace(self) -> None: + self._prev_trace = sys.gettrace() sys.settrace(self._trace_func) threading.settrace(self._trace_func) def _stop_settrace(self) -> None: - sys.settrace(None) - threading.settrace(None) + sys.settrace(self._prev_trace) + threading.settrace(None) # no public API to get previous def _trace_func(self, frame: Any, event: str, arg: Any) -> Any: if event == "call": @@ -653,34 +682,32 @@ def _extract_rows( limit: int | None = None, ) -> list[dict[str, Any]]: self._ensure_stopped() - sort_key = ( - Profiler._resolve_sort_key(sort_by) if sort_by else self._default_sort - ) + sort_key = _resolve_sort_key(sort_by) if sort_by else self._default_sort total_wall = self.total_time or 1e-9 records = [r for r in self._records if r.file != _SELF_FILE] - by_thread: dict[int, list[TraceRecord]] = defaultdict(list) - for r in records: - by_thread[r.thread_id].append(r) + by_thread: dict[int, list[tuple[int, TraceRecord]]] = defaultdict(list) + for i, r in enumerate(records): + by_thread[r.thread_id].append((i, r)) child_sum: dict[int, float] = defaultdict(float) - for _tid, recs in by_thread.items(): - recs.sort(key=lambda r: r.start_ns) + for _tid, indexed_recs in by_thread.items(): + indexed_recs.sort(key=lambda x: x[1].start_ns) parent_stack: list[tuple[int, TraceRecord]] = [] - for i, rec in enumerate(recs): + for idx, rec in indexed_recs: while parent_stack and parent_stack[-1][1].end_ns <= rec.start_ns: parent_stack.pop() if parent_stack and parent_stack[-1][1].depth == rec.depth - 1: parent_idx = parent_stack[-1][0] child_sum[parent_idx] += (rec.end_ns - rec.start_ns) / 1e9 - parent_stack.append((id(rec), rec)) + parent_stack.append((idx, rec)) agg: dict[tuple[str, str, int], dict[str, Any]] = {} - for rec in records: + for idx, rec in enumerate(records): key = (rec.func, rec.file, rec.lineno) duration = (rec.end_ns - rec.start_ns) / 1e9 - self_time = max(duration - child_sum.get(id(rec), 0.0), 0.0) + self_time = max(duration - child_sum.get(idx, 0.0), 0.0) if key not in agg: agg[key] = { "func": rec.func, @@ -776,9 +803,9 @@ def _compute_self_time(node: dict[str, Any]) -> None: _compute_self_time(child) def _merge_children(node: dict[str, Any]) -> None: - merged: dict[str, dict[str, Any]] = {} + merged: dict[tuple[str, str, int], dict[str, Any]] = {} for child in node["children"]: - key = child["name"] + key = (child["name"], child["file"], child["lineno"]) if key in merged: merged[key]["cumtime"] += child["cumtime"] merged[key]["tottime"] += child["tottime"] @@ -824,9 +851,7 @@ def output_text( total_calls = sum(r["calls"] for r in rows) total_time_s = self.total_time - sort_key = ( - Profiler._resolve_sort_key(sort_by) if sort_by else self._default_sort - ) + sort_key = _resolve_sort_key(sort_by) if sort_by else self._default_sort lines: list[str] = [ f" {total_calls} function calls in {total_time_s:.3f} seconds\n", diff --git a/profiler/test_profiler_correctness.py b/profiler/test_profiler_correctness.py index b128e70..3bb7cc5 100644 --- a/profiler/test_profiler_correctness.py +++ b/profiler/test_profiler_correctness.py @@ -651,16 +651,6 @@ def test_icicle_to_file(self): assert f.read() == html -# -- Yappi availability flag ------------------------------------------------- - -try: - import yappi as _yappi - - _HAS_YAPPI = True -except ImportError: - _HAS_YAPPI = False - - # ============================================================================ # TracingProfiler tests # ============================================================================ @@ -736,6 +726,10 @@ def test_context_manager_returns_self(self): _busy_work() assert p is tp + def test_builtins_raises_not_implemented(self): + with pytest.raises(NotImplementedError, match="C-level"): + TracingProfiler(builtins=True) + # -- TestTracingAsyncProfiler ------------------------------------------------ @@ -1198,6 +1192,10 @@ def test_very_short_execution(self): text = p.output_text() assert isinstance(text, str) + def test_builtins_raises_not_implemented(self): + with pytest.raises(NotImplementedError, match="C-level"): + TracingProfiler(builtins=True) + # -- TestTracingVsYappi ------------------------------------------------------ From b812c770a2eea66fb8a50bb4433a6a6c2a575eda Mon Sep 17 00:00:00 2001 From: Peng Ding Date: Tue, 15 Sep 2026 03:31:09 -0500 Subject: [PATCH 3/3] fix: merge flamegraph roots for yielding coroutines, add GIL note - Merge root-level nodes by (name, file, lineno) in _extract_call_tree so coroutines split by yield/resume appear as a single flamegraph entry - Add comment noting _stacks defaultdict is lock-free under GIL and needs locking for free-threaded builds --- profiler/profiler.py | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/profiler/profiler.py b/profiler/profiler.py index 473656e..72a99da 100644 --- a/profiler/profiler.py +++ b/profiler/profiler.py @@ -487,6 +487,7 @@ def __init__( self._running = False self._records: list[TraceRecord] = [] self._records_lock = threading.Lock() + # Per-thread stacks — lock-free under GIL; needs lock for free-threaded builds self._stacks: dict[int, list[tuple[str, str, int, int]]] = defaultdict(list) self._tool_id: int | None = None self._wall_start_ns: int = 0 @@ -824,7 +825,23 @@ def _merge_children(node: dict[str, Any]) -> None: _compute_self_time(root) _merge_children(root) - all_roots.sort(key=lambda n: n["cumtime"], reverse=True) + merged_roots: dict[tuple[str, str, int], dict[str, Any]] = {} + for root in all_roots: + key = (root["name"], root["file"], root["lineno"]) + if key in merged_roots: + merged_roots[key]["cumtime"] += root["cumtime"] + merged_roots[key]["tottime"] += root["tottime"] + merged_roots[key]["calls"] += root["calls"] + merged_roots[key]["cumtime_pct"] += root["cumtime_pct"] + merged_roots[key]["children"].extend(root["children"]) + else: + merged_roots[key] = root + all_roots = sorted( + merged_roots.values(), key=lambda n: n["cumtime"], reverse=True + ) + for root in all_roots: + _merge_children(root) + return all_roots # -- Output methods ------------------------------------------------------