"""Run-local timing of Python orchestration; numeric timings come from C.""" from __future__ import annotations from collections.abc import Callable, Generator, Mapping from contextlib import contextmanager from contextvars import ContextVar from dataclasses import dataclass, field from functools import wraps import inspect import os from time import perf_counter_ns from typing import Any, Literal, ParamSpec, TypeVar, cast ProfileMode = Literal['off', 'standard', 'audit'] _P = ParamSpec('_P') _R = TypeVar('_R') _MODE_RANK: Mapping[ProfileMode, int] = {'off': 0, 'standard': 1, 'audit': 2} def _read_startup_mode() -> ProfileMode: raw_mode = os.getenv("SIMULATIONAPP_PROFILE", "off").strip().lower() aliases: dict[str, ProfileMode] = { "": "off", "0": "off", "false": "off", "no": "off", "off": "off", "1": "standard", "true": "standard", "yes": "standard", "on": "standard", "standard": "standard", "audit": "audit", } try: return aliases[raw_mode] except KeyError as exc: raise ValueError( "SIMULATIONAPP_PROFILE must be one of: off, standard, audit." ) from exc def _minimum_mode(value: str) -> ProfileMode: normalized = value.strip().lower() if normalized not in _MODE_RANK: raise ValueError("minimum_mode must be one of: off, standard, audit.") return cast(ProfileMode, normalized) def _mode_enabled(minimum_mode: ProfileMode) -> bool: # ``off`` is an unconditional zero-wrapper mode, even if a caller passes # ``minimum_mode="off"`` by mistake. return ( PROFILE_MODE != "off" and _MODE_RANK[PROFILE_MODE] >= _MODE_RANK[minimum_mode] ) @dataclass class _TimingStats: calls: int = 0 inclusive_ns: int = 0 self_ns: int = 0 max_ns: int = 0 errors: int = 0 def record(self, inclusive_ns: int, self_ns: int, error: bool) -> None: self.calls += 1 self.inclusive_ns += inclusive_ns self.self_ns += self_ns self.max_ns = max(self.max_ns, inclusive_ns) if error: self.errors += 1 def snapshot(self) -> dict[str, int]: return { "calls": self.calls, "inclusiveNs": self.inclusive_ns, "selfNs": self.self_ns, "maxNs": self.max_ns, "errors": self.errors, } def profile_phase( name: str, minimum_mode: str = "standard", ) -> Callable[[Callable[_P, _R]], Callable[_P, _R]]: """Decorate a simulation phase while preserving the off-mode callable.""" minimum = _minimum_mode(minimum_mode) def decorate(function: Callable[_P, _R]) -> Callable[_P, _R]: if not _mode_enabled(minimum): return function if inspect.iscoroutinefunction(function): @wraps(function) async def async_wrapper(*args: _P.args, **kwargs: _P.kwargs) -> Any: trace = _CURRENT_TRACE.get() if trace is None: return await function(*args, **kwargs) with _tracked_span( trace, name, ): return await function(*args, **kwargs) return cast(Callable[_P, _R], async_wrapper) @wraps(function) def wrapper(*args: _P.args, **kwargs: _P.kwargs) -> _R: trace = _CURRENT_TRACE.get() if trace is None: return function(*args, **kwargs) with _tracked_span( trace, name, ): return function(*args, **kwargs) return wrapper return decorate PROFILE_MODE = _read_startup_mode() @dataclass class _ActiveSpan: name: str started_ns: int child_ns: int = 0 @dataclass class PerformanceTrace: mode: ProfileMode _phases: dict[str, _TimingStats] = field(default_factory=dict, repr=False) @property def enabled(self): return self.mode != 'off' def snapshot(self): return {'mode': self.mode, 'phases': {key: value.snapshot() for key, value in sorted(self._phases.items())}} _CURRENT_TRACE = ContextVar('simulation_trace', default=None) _ACTIVE_SPANS = ContextVar('simulation_spans', default=()) @contextmanager def _tracked_span(trace, name): stack = _ACTIVE_SPANS.get() frame = _ActiveSpan(name, perf_counter_ns()) token = _ACTIVE_SPANS.set((*stack, frame)) error = False try: yield except BaseException: error = True raise finally: elapsed = max(0, perf_counter_ns()-frame.started_ns) _ACTIVE_SPANS.reset(token) if stack: stack[-1].child_ns += elapsed trace._phases.setdefault(name, _TimingStats()).record(elapsed, max(0, elapsed-frame.child_ns), error) @contextmanager def profile_run(): trace = PerformanceTrace(PROFILE_MODE) token = _CURRENT_TRACE.set(trace) spans = _ACTIVE_SPANS.set(()) try: if trace.enabled: with _tracked_span(trace, 'simulation.total'): yield trace else: yield trace finally: _ACTIVE_SPANS.reset(spans) _CURRENT_TRACE.reset(token) @contextmanager def performance_span(name: str, minimum_mode: str = 'standard'): minimum = _minimum_mode(minimum_mode) trace = _CURRENT_TRACE.get() if trace is None or not _mode_enabled(minimum): yield return with _tracked_span(trace, name): yield