Skip to main content

candle_graph/instrument/
session.rs

1//! Representative-run profiler session — emits `candle-graph/trace/10` JSONL.
2
3use std::cell::RefCell;
4use std::collections::BTreeMap;
5use std::fs::{File, OpenOptions};
6use std::io::{self, Write};
7use std::path::{Path, PathBuf};
8use std::time::{Instant, SystemTime, UNIX_EPOCH};
9
10use anyhow::{Context, Result};
11use serde::Serialize;
12
13use crate::capability::CaptureContract;
14use crate::phase::ExecutionPhase;
15use crate::trace::events::{
16    DeviceIntervalEvent, DeviceMemoryEvent, EdgeEvent, GradientEvent, MemoryEvent, OpEvent,
17    SpanEndEvent, SpanStartEvent, TensorEvent, TensorStatsEvent, TerminalEvent, TraceEvent,
18};
19use crate::trace::memory::{resolve_dense_tensor_bytes, MemoryAction};
20use crate::trace::schema::{
21    ComparisonIdentity, GradientState, RunOutcome, TimingMode, TraceRunMeta,
22};
23
24use super::span::{
25    DeviceIntervalRecord, DeviceMemoryRecord, MemoryRecord, OpRecord, SpanGuard, SpanId, SpanKind,
26    TensorRecord,
27};
28
29/// Streaming trace session writing TensorFlow-Profiler-style span JSONL.
30pub struct TraceSession {
31    path: PathBuf,
32    inner: RefCell<SessionInner>,
33}
34
35/// Required provenance for one representative profile run.
36#[derive(Debug, Clone, PartialEq, Eq)]
37pub struct ProfileRun {
38    pub entrypoint: String,
39    pub correlation_id: String,
40    pub phase: ExecutionPhase,
41    /// One-based selected update or inference invocation.
42    pub capture_step: u64,
43    pub warmup_steps: u64,
44    pub device: String,
45    pub measured_region_device_synchronized: bool,
46    pub timing_mode: TimingMode,
47    pub capture_contract: CaptureContract,
48    pub comparison_identity: Option<ComparisonIdentity>,
49    pub tags: BTreeMap<String, String>,
50}
51
52impl ProfileRun {
53    pub fn training(
54        entrypoint: impl Into<String>,
55        capture_step: u64,
56        device: impl Into<String>,
57    ) -> Self {
58        let entrypoint = entrypoint.into();
59        Self {
60            correlation_id: format!("{entrypoint}/update-{capture_step}"),
61            entrypoint,
62            phase: ExecutionPhase::Train,
63            capture_step,
64            warmup_steps: capture_step.saturating_sub(1),
65            device: device.into(),
66            measured_region_device_synchronized: false,
67            timing_mode: TimingMode::Host,
68            capture_contract: CaptureContract::default(),
69            comparison_identity: None,
70            tags: BTreeMap::new(),
71        }
72    }
73
74    pub fn tag(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
75        self.tags.insert(key.into(), value.into());
76        self
77    }
78
79    pub fn correlation_id(mut self, value: impl Into<String>) -> Self {
80        self.correlation_id = value.into();
81        self
82    }
83
84    pub fn device_synchronized(mut self) -> Self {
85        self.timing_mode = TimingMode::DeviceSynchronized;
86        self.measured_region_device_synchronized = true;
87        self
88    }
89
90    /// Mark only the caller-controlled measured region as device-synchronized.
91    /// Nested semantic spans remain host-timed.
92    pub fn measured_region_device_synchronized(mut self) -> Self {
93        self.measured_region_device_synchronized = true;
94        self
95    }
96
97    pub fn inference(
98        entrypoint: impl Into<String>,
99        capture_step: u64,
100        device: impl Into<String>,
101    ) -> Self {
102        let entrypoint = entrypoint.into();
103        Self {
104            correlation_id: format!("{entrypoint}/inference-{capture_step}"),
105            entrypoint,
106            phase: ExecutionPhase::Infer,
107            capture_step,
108            warmup_steps: capture_step.saturating_sub(1),
109            device: device.into(),
110            measured_region_device_synchronized: false,
111            timing_mode: TimingMode::Host,
112            capture_contract: CaptureContract::default(),
113            comparison_identity: None,
114            tags: BTreeMap::new(),
115        }
116    }
117
118    pub fn capture_contract(mut self, contract: CaptureContract) -> Self {
119        self.capture_contract = contract;
120        self
121    }
122
123    pub fn comparison_identity(mut self, identity: ComparisonIdentity) -> Self {
124        self.comparison_identity = Some(identity);
125        self
126    }
127}
128
129struct SessionInner {
130    writer: io::BufWriter<File>,
131    span_stack: Vec<u64>,
132    next_span_id: u64,
133    next_event_id: u64,
134    id_buf: String,
135    probe_started: Instant,
136    sticky_error: Option<String>,
137}
138
139impl SessionInner {
140    /// Write one event, failing closed on any prior failure.
141    ///
142    /// Invariant: any write failure poisons the session via `sticky_error`, so
143    /// a trace whose stream may be corrupt or incomplete can never be finished
144    /// as `Complete`.
145    fn write(&mut self, event: &TraceEvent) -> Result<()> {
146        if let Some(error) = &self.sticky_error {
147            anyhow::bail!("trace session previously failed: {error}");
148        }
149        if let Err(error) = write_event(&mut self.writer, event) {
150            self.sticky_error.get_or_insert_with(|| error.to_string());
151            return Err(error);
152        }
153        Ok(())
154    }
155}
156
157impl TraceSession {
158    /// Open a trace and own its single root span until [`Self::finish`].
159    pub fn open(path: impl AsRef<Path>, run: ProfileRun) -> Result<Self> {
160        let entrypoint = run.entrypoint.clone();
161        let meta = TraceRunMeta {
162            run_id: new_run_id(),
163            correlation_id: run.correlation_id,
164            entrypoint: run.entrypoint,
165            phase: run.phase,
166            timestamp: utc_iso8601_now(),
167            capture_step: run.capture_step,
168            warmup_steps: run.warmup_steps,
169            device: run.device,
170            measured_region_device_synchronized: run.measured_region_device_synchronized,
171            timing_mode: run.timing_mode,
172            capture_contract: run.capture_contract,
173            comparison_identity: run.comparison_identity,
174            tags: run.tags,
175            candle_version: None,
176        };
177        meta.validate().context("validate run provenance")?;
178        meta.capture_contract
179            .validate()
180            .context("validate capture contract")?;
181        let path = path.as_ref().to_path_buf();
182        if let Some(parent) = path.parent() {
183            std::fs::create_dir_all(parent)
184                .with_context(|| format!("create trace dir {}", parent.display()))?;
185        }
186        let file = OpenOptions::new()
187            .create(true)
188            .write(true)
189            .truncate(true)
190            .open(&path)
191            .with_context(|| format!("open trace {}", path.display()))?;
192        let mut writer = io::BufWriter::new(file);
193        write_event(&mut writer, &TraceEvent::meta(meta))?;
194        write_event(
195            &mut writer,
196            &TraceEvent::SpanStart(SpanStartEvent {
197                id: span_id_string(1),
198                parent_id: None,
199                name: entrypoint,
200                start_ns: 0,
201                kind: SpanKind::Function,
202                measured: false,
203                step: None,
204            }),
205        )?;
206        Ok(Self {
207            path,
208            inner: RefCell::new(SessionInner {
209                writer,
210                span_stack: vec![1],
211                next_span_id: 1,
212                next_event_id: 0,
213                id_buf: String::with_capacity(24),
214                probe_started: Instant::now(),
215                sticky_error: None,
216            }),
217        })
218    }
219
220    /// Begin a nested span; parent is the top of the session span stack (TF Profiler call tree).
221    pub fn begin_span(&self, name: impl Into<String>, kind: SpanKind) -> SpanGuard<'_> {
222        self.begin_span_inner(name, kind, None, false)
223    }
224
225    /// Begin the single caller-controlled region used for total-time comparisons.
226    pub fn begin_measurement(&self, name: impl Into<String>) -> SpanGuard<'_> {
227        self.begin_span_inner(name, SpanKind::Function, None, true)
228    }
229
230    /// Begin a span tagged with a PyTorch-style training step (`forward` / `backward` / `optimizer`).
231    pub fn begin_step_span(
232        &self,
233        name: impl Into<String>,
234        step: crate::phase::ExecutionStep,
235        kind: SpanKind,
236    ) -> SpanGuard<'_> {
237        self.begin_span_inner(name, kind, Some(step), false)
238    }
239
240    /// Record already-completed host work without changing the live span stack.
241    ///
242    /// `start_ns` is a monotonic offset from this session's start and `duration_ns` must be
243    /// positive. The interval must have completed before this call. The closed span is attached
244    /// directly to the session root, so it may overlap live nested spans without implying a
245    /// synchronous parent/child call relationship.
246    pub fn record_completed_host_span(
247        &self,
248        name: impl Into<String>,
249        kind: SpanKind,
250        start_ns: u64,
251        duration_ns: u64,
252    ) -> Result<SpanId> {
253        anyhow::ensure!(
254            duration_ns > 0,
255            "completed host span duration must be positive"
256        );
257        let end_ns = start_ns
258            .checked_add(duration_ns)
259            .context("completed host span interval overflows u64 nanoseconds")?;
260        anyhow::ensure!(
261            end_ns <= self.elapsed_ns(),
262            "completed host span ends in the future relative to the trace session"
263        );
264
265        let mut inner = self.inner.borrow_mut();
266        if let Some(error) = &inner.sticky_error {
267            anyhow::bail!("trace session previously failed: {error}");
268        }
269        let root_id = inner
270            .span_stack
271            .first()
272            .copied()
273            .context("trace session root span is missing")?;
274        inner.next_span_id += 1;
275        let span_id = SpanId(inner.next_span_id);
276        let id = span_id_string(span_id.0);
277
278        inner.write(&TraceEvent::SpanStart(SpanStartEvent {
279            id: id.clone(),
280            parent_id: Some(span_id_string(root_id)),
281            name: name.into(),
282            start_ns,
283            kind,
284            measured: false,
285            step: None,
286        }))?;
287        inner.write(&TraceEvent::SpanEnd(SpanEndEvent { id, duration_ns }))?;
288        Ok(span_id)
289    }
290
291    fn begin_span_inner(
292        &self,
293        name: impl Into<String>,
294        kind: SpanKind,
295        step: Option<crate::phase::ExecutionStep>,
296        measured: bool,
297    ) -> SpanGuard<'_> {
298        let started = Instant::now();
299        let start_ns = self.elapsed_ns();
300        let mut inner = self.inner.borrow_mut();
301        inner.next_span_id += 1;
302        let span_id = inner.next_span_id;
303        let parent_id = inner.span_stack.last().copied().map(span_id_string);
304
305        format_span_id(&mut inner.id_buf, span_id);
306        let id_str = inner.id_buf.clone();
307
308        // The guard must still be returned on failure; the helper has already
309        // poisoned the session, so the error is safe to drop here.
310        let _ = inner.write(&TraceEvent::SpanStart(SpanStartEvent {
311            id: id_str,
312            parent_id,
313            name: name.into(),
314            start_ns,
315            kind,
316            measured,
317            step,
318        }));
319
320        inner.span_stack.push(span_id);
321
322        SpanGuard {
323            session: self,
324            id: SpanId(span_id),
325            started,
326        }
327    }
328
329    pub(crate) fn end_span(&self, id: SpanId, duration_ns: u64) -> Result<()> {
330        let mut inner = self.inner.borrow_mut();
331        let expected = inner
332            .span_stack
333            .last()
334            .copied()
335            .with_context(|| format!("span stack underflow closing span {}", id.0))?;
336        anyhow::ensure!(
337            expected == id.0,
338            "span_end id `{}` does not match open span `{}`",
339            id.0,
340            expected
341        );
342        inner.span_stack.pop();
343
344        format_span_id(&mut inner.id_buf, id.0);
345        let span_id = inner.id_buf.clone();
346        inner.write(&TraceEvent::SpanEnd(SpanEndEvent {
347            id: span_id,
348            duration_ns,
349        }))
350    }
351
352    pub fn elapsed_ns(&self) -> u64 {
353        self.inner
354            .borrow()
355            .probe_started
356            .elapsed()
357            .as_nanos()
358            .min(u64::MAX as u128) as u64
359    }
360
361    /// Convert an [`Instant`] from this process into this trace session's
362    /// monotonic nanosecond clock without an independently sampled anchor.
363    pub fn host_timestamp_ns(&self, instant: Instant) -> Result<u64> {
364        let probe_started = self.inner.borrow().probe_started;
365        let elapsed = instant
366            .checked_duration_since(probe_started)
367            .context("host instant predates the trace session")?;
368        Ok(elapsed.as_nanos().min(u64::MAX as u128) as u64)
369    }
370
371    /// Record a timed op observation attached to `span_id`.
372    pub fn record_op(&self, span_id: SpanId, op: OpRecord<'_>) -> Result<()> {
373        let output_dense_bytes =
374            resolve_dense_tensor_bytes(op.output_dense_bytes, op.shape, op.dtype);
375        let timestamp_ns = if op.timestamp_ns > 0 {
376            op.timestamp_ns
377        } else {
378            self.elapsed_ns().saturating_sub(op.duration_ns)
379        };
380        {
381            let mut inner = self.inner.borrow_mut();
382            inner.write(&TraceEvent::Op(OpEvent {
383                span_id: span_id_string(span_id.0),
384                op_name: op.op_name.into(),
385                inputs: op.inputs.to_vec(),
386                output: op.output.map(str::to_string),
387                shape: op.shape.to_vec(),
388                dtype: op.dtype.into(),
389                device: op.device.into(),
390                duration_ns: op.duration_ns,
391                timestamp_ns,
392                output_dense_bytes,
393                input_dense_bytes: op.input_dense_bytes,
394            }))?;
395        }
396
397        Ok(())
398    }
399
400    /// Record tensor metadata. Logical allocation lifetimes require explicit memory events.
401    pub fn record_tensor(&self, span_id: SpanId, tensor: TensorRecord<'_>) -> Result<()> {
402        let dense_bytes =
403            resolve_dense_tensor_bytes(tensor.dense_bytes, tensor.shape, tensor.dtype);
404        let mut inner = self.inner.borrow_mut();
405        inner.write(&TraceEvent::Tensor(TensorEvent {
406            span_id: span_id_string(span_id.0),
407            tensor_id: tensor.tensor_id.into(),
408            label: tensor.label.map(str::to_string),
409            shape: tensor.shape.to_vec(),
410            dtype: tensor.dtype.into(),
411            device: tensor.device.into(),
412            requires_grad: tensor.requires_grad,
413            dense_bytes,
414            category: tensor.category,
415        }))?;
416        Ok(())
417    }
418
419    /// Record one application-computed scalar observation (loss term, gate value,
420    /// clip fraction, cosine, …) already available on the host.
421    ///
422    /// No device reductions or readbacks are triggered; the value is stored as a
423    /// single-element labeled statistic on the tensor-stats plane, so series and
424    /// comparison tooling treat it like any other labeled mechanism observation.
425    /// The evidence is observational; it never claims causal attribution.
426    pub fn record_scalar(&self, span_id: SpanId, label: &str, value: f64) -> Result<()> {
427        anyhow::ensure!(!label.trim().is_empty(), "scalar labels must not be empty");
428        let finite = value.is_finite();
429        // JSON has no NaN/inf representation; the non_finite count is
430        // authoritative and the numeric fields stay serializable.
431        let event = TensorStatsEvent {
432            span_id: span_id_string(span_id.0),
433            label: label.to_string(),
434            shape: Vec::new(),
435            dtype: "f64".into(),
436            elements: 1,
437            non_finite: u64::from(!finite),
438            rms: if finite { value.abs() } else { 0.0 },
439            abs_max: if finite { value.abs() } else { 0.0 },
440            mean: if finite { value } else { 0.0 },
441        };
442        let mut inner = self.inner.borrow_mut();
443        inner.write(&TraceEvent::TensorStats(event))
444    }
445
446    /// Record device-reduced numerical statistics for a caller-labeled Candle tensor.
447    #[cfg(feature = "candle")]
448    pub fn record_tensor_stats(
449        &self,
450        span_id: SpanId,
451        label: &str,
452        tensor: &candle_core::Tensor,
453    ) -> Result<()> {
454        use candle_core::DType;
455
456        let elements = tensor.elem_count() as u64;
457        let (non_finite, rms, abs_max, mean) = if elements == 0 {
458            (0, 0.0, 0.0, 0.0)
459        } else {
460            let x = tensor
461                .detach()
462                .to_dtype(DType::F32)
463                .with_context(|| format!("cast tensor stats {label:?} to f32"))?;
464            let finite = x
465                .sub(&x)
466                .and_then(|delta| delta.eq(0.0))
467                .and_then(|mask| mask.to_dtype(DType::U32))
468                .and_then(|mask| mask.sum_all())
469                .and_then(|count| count.to_scalar::<u32>())
470                .with_context(|| format!("count finite elements for tensor stats {label:?}"))?
471                as u64;
472            let read = |value: candle_core::Result<candle_core::Tensor>, name: &str| {
473                value
474                    .and_then(|value| value.to_scalar::<f32>())
475                    .map(f64::from)
476                    .with_context(|| format!("reduce {name} for tensor stats {label:?}"))
477            };
478            let rms = read(x.sqr().and_then(|x2| x2.mean_all()?.sqrt()), "rms")?;
479            let abs_max = read(x.abs().and_then(|abs| abs.max_all()), "abs_max")?;
480            let mean = read(x.mean_all(), "mean")?;
481            // JSON has no NaN/inf number representation. The explicit
482            // non-finite count is authoritative; keep the remaining fields
483            // serializable when an all-element reduction is non-finite.
484            let json_number = |value: f64| if value.is_finite() { value } else { 0.0 };
485            (
486                elements.saturating_sub(finite),
487                json_number(rms),
488                json_number(abs_max),
489                json_number(mean),
490            )
491        };
492        let event = TensorStatsEvent {
493            span_id: span_id_string(span_id.0),
494            label: label.to_string(),
495            shape: tensor.dims().to_vec(),
496            dtype: format!("{:?}", tensor.dtype()).to_ascii_lowercase(),
497            elements,
498            non_finite,
499            rms,
500            abs_max,
501            mean,
502        };
503        let mut inner = self.inner.borrow_mut();
504        inner.write(&TraceEvent::TensorStats(event))
505    }
506
507    /// Record an explicit tensor allocation (TensorFlow memory timeline).
508    pub fn record_memory_alloc(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
509        let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
510        let mut inner = self.inner.borrow_mut();
511        inner.write(&TraceEvent::Memory(MemoryEvent {
512            timestamp_ns,
513            storage_id: mem.storage_id.into(),
514            tensor_id: mem.tensor_id.into(),
515            span_id: span_id_string(span_id.0),
516            op_name: mem.op_name.map(str::to_string),
517            device: mem.device.into(),
518            bytes: mem.bytes,
519            action: MemoryAction::Alloc,
520            shape: mem.shape.to_vec(),
521            dtype: mem.dtype.into(),
522            category: mem.category,
523        }))
524    }
525
526    /// Record an explicit tensor deallocation.
527    pub fn record_memory_free(&self, span_id: SpanId, mem: MemoryRecord<'_>) -> Result<()> {
528        let timestamp_ns = mem.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
529        let mut inner = self.inner.borrow_mut();
530        inner.write(&TraceEvent::Memory(MemoryEvent {
531            timestamp_ns,
532            storage_id: mem.storage_id.into(),
533            tensor_id: mem.tensor_id.into(),
534            span_id: span_id_string(span_id.0),
535            op_name: mem.op_name.map(str::to_string),
536            device: mem.device.into(),
537            bytes: mem.bytes,
538            action: MemoryAction::Free,
539            shape: mem.shape.to_vec(),
540            dtype: mem.dtype.into(),
541            category: mem.category,
542        }))
543    }
544
545    /// Record a device-level memory checkpoint (cudaMemGetInfo-style).
546    pub fn record_device_memory(&self, sample: DeviceMemoryRecord<'_>) -> Result<()> {
547        anyhow::ensure!(
548            sample.used_bytes.is_some()
549                || sample.free_bytes.is_some()
550                || sample.reserved_bytes.is_some()
551                || sample.capacity_bytes.is_some(),
552            "device-memory checkpoint must contain at least one observation"
553        );
554        let timestamp_ns = sample.timestamp_ns.unwrap_or_else(|| self.elapsed_ns());
555        let mut inner = self.inner.borrow_mut();
556        inner.write(&TraceEvent::DeviceMemory(DeviceMemoryEvent {
557            timestamp_ns,
558            device: sample.device.into(),
559            used_bytes: sample.used_bytes,
560            free_bytes: sample.free_bytes,
561            reserved_bytes: sample.reserved_bytes,
562            capacity_bytes: sample.capacity_bytes,
563        }))
564    }
565
566    pub fn record_device_interval(
567        &self,
568        span_id: SpanId,
569        interval: DeviceIntervalRecord<'_>,
570    ) -> Result<()> {
571        anyhow::ensure!(
572            interval.duration_ns > 0,
573            "device interval duration must be positive"
574        );
575        let mut inner = self.inner.borrow_mut();
576        inner.write(&TraceEvent::DeviceInterval(DeviceIntervalEvent {
577            span_id: span_id_string(span_id.0),
578            device: interval.device.into(),
579            stream_id: interval.stream_id.into(),
580            clock_id: interval.clock_id.into(),
581            backend: interval.backend.into(),
582            start_ns: interval.start_ns,
583            duration_ns: interval.duration_ns,
584        }))
585    }
586
587    pub fn record_call_edge(&self, from: SpanId, to: SpanId, duration_ns: u64) -> Result<()> {
588        let mut inner = self.inner.borrow_mut();
589        inner.write(&TraceEvent::Edge(EdgeEvent::Call {
590            from_span: span_id_string(from.0),
591            to_span: span_id_string(to.0),
592            host_duration_ns: duration_ns,
593        }))
594    }
595
596    pub fn record_data_edge(&self, from_tensor: &str, to_tensor: &str) -> Result<()> {
597        anyhow::ensure!(
598            !from_tensor.is_empty() && !to_tensor.is_empty(),
599            "data-edge tensor IDs cannot be empty"
600        );
601        let mut inner = self.inner.borrow_mut();
602        inner.write(&TraceEvent::Edge(EdgeEvent::Data {
603            from_tensor: from_tensor.into(),
604            to_tensor: to_tensor.into(),
605        }))
606    }
607
608    /// Record one parameter gradient fact from a probe run.
609    ///
610    /// `Present` requires a finite positive norm, `Zero` requires positive zero, and `Missing` or
611    /// `NonFinite` require `None`. Exact-contract captures emit one event per `(root, key)`.
612    pub fn record_gradient(
613        &self,
614        root: impl Into<String>,
615        key: impl Into<String>,
616        state: GradientState,
617        norm: Option<f64>,
618    ) -> Result<()> {
619        let root = root.into();
620        let key = key.into();
621        anyhow::ensure!(
622            !root.trim().is_empty() && !key.trim().is_empty(),
623            "gradient roots and parameter keys must not be empty"
624        );
625        anyhow::ensure!(
626            state.norm_is_valid(norm),
627            "gradient state `{state}` is inconsistent with norm {norm:?}"
628        );
629        let mut inner = self.inner.borrow_mut();
630        inner.next_event_id += 1;
631        let event_id = format!("gradient-{}", inner.next_event_id);
632        inner.write(&TraceEvent::Gradient(GradientEvent {
633            event_id,
634            root,
635            key,
636            state,
637            norm,
638        }))
639    }
640
641    pub fn flush(&self) -> Result<()> {
642        let mut inner = self.inner.borrow_mut();
643        if let Some(error) = &inner.sticky_error {
644            anyhow::bail!("trace session previously failed: {error}");
645        }
646        inner.writer.flush().context("flushing trace JSONL")
647    }
648
649    /// Close the owned root span, flush, and return the trace path.
650    pub fn finish(self) -> Result<PathBuf> {
651        let duration_ns = self.elapsed_ns();
652        {
653            let mut inner = self.inner.borrow_mut();
654            if let Some(error) = &inner.sticky_error {
655                anyhow::bail!("trace session previously failed: {error}");
656            }
657            anyhow::ensure!(
658                inner.span_stack.as_slice() == [1],
659                "cannot finish trace with {} nested spans still open",
660                inner.span_stack.len().saturating_sub(1)
661            );
662            inner.span_stack.pop();
663            inner.write(&TraceEvent::SpanEnd(SpanEndEvent {
664                id: span_id_string(1),
665                duration_ns,
666            }))?;
667            inner.write(&TraceEvent::Terminal(TerminalEvent {
668                outcome: RunOutcome::Complete,
669                timestamp_ns: duration_ns,
670                reason: None,
671            }))?;
672            inner.writer.flush().context("flushing trace JSONL")?;
673        }
674        Ok(self.path)
675    }
676
677    /// Finalize a diagnosable partial trace without presenting it as complete evidence.
678    pub fn finish_failed(self, reason: impl Into<String>) -> Result<PathBuf> {
679        let reason = reason.into();
680        anyhow::ensure!(!reason.trim().is_empty(), "failure reason cannot be empty");
681        let timestamp_ns = self.elapsed_ns();
682        let mut inner = self.inner.borrow_mut();
683        write_event(
684            &mut inner.writer,
685            &TraceEvent::Terminal(TerminalEvent {
686                outcome: RunOutcome::Failed,
687                timestamp_ns,
688                reason: Some(reason),
689            }),
690        )?;
691        inner
692            .writer
693            .flush()
694            .context("flushing failed trace JSONL")?;
695        drop(inner);
696        Ok(self.path.clone())
697    }
698}
699
700fn write_event<W: Write, T: Serialize>(writer: &mut W, event: &T) -> Result<()> {
701    let mut line = serde_json::to_vec(event).context("serializing trace JSONL event")?;
702    line.push(b'\n');
703    writer.write_all(&line).context("writing trace JSONL event")
704}
705
706fn format_span_id(buf: &mut String, id: u64) {
707    buf.clear();
708    use std::fmt::Write as _;
709    let _ = write!(buf, "s{id}");
710}
711
712fn span_id_string(id: u64) -> String {
713    format!("s{id}")
714}
715
716fn new_run_id() -> String {
717    let pid = std::process::id();
718    let nanos = SystemTime::now()
719        .duration_since(UNIX_EPOCH)
720        .map(|d| d.as_nanos())
721        .unwrap_or(0);
722    format!("run-{pid}-{nanos}")
723}
724
725fn utc_iso8601_now() -> String {
726    let now = SystemTime::now()
727        .duration_since(UNIX_EPOCH)
728        .expect("system clock before UNIX epoch");
729    let secs = now.as_secs();
730    let millis = now.subsec_millis();
731    let (year, month, day, hour, minute, second) = unix_secs_to_utc(secs);
732    format!("{year:04}-{month:02}-{day:02}T{hour:02}:{minute:02}:{second:02}.{millis:03}Z")
733}
734
735/// Convert Unix seconds to UTC calendar components (Gregorian).
736fn unix_secs_to_utc(secs: u64) -> (u64, u64, u64, u64, u64, u64) {
737    const SECS_PER_DAY: u64 = 86_400;
738    let days = secs / SECS_PER_DAY;
739    let rem = secs % SECS_PER_DAY;
740    let hour = rem / 3600;
741    let minute = (rem % 3600) / 60;
742    let second = rem % 60;
743
744    let z = days + 719_468;
745    let era = z / 146_097;
746    let doe = z - era * 146_097;
747    let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365;
748    let y = yoe + era * 400;
749    let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
750    let mp = (5 * doy + 2) / 153;
751    let day = doy - (153 * mp + 2) / 5 + 1;
752    let month = if mp < 10 { mp + 3 } else { mp - 9 };
753    let year = if month <= 2 { y + 1 } else { y };
754    (year, month, day, hour, minute, second)
755}
756
757#[cfg(test)]
758mod tests {
759    use super::*;
760    use crate::trace::document::parse_trace;
761    use crate::trace::events::SpanStartEvent;
762    use crate::trace::health::analyze_health;
763    use serde_json::Value;
764    use std::io::{BufRead, BufReader};
765
766    fn temp_trace(name: &str) -> PathBuf {
767        std::env::temp_dir().join(format!(
768            "candle-graph-trace-{}-{}-{name}",
769            std::process::id(),
770            std::time::SystemTime::now()
771                .duration_since(UNIX_EPOCH)
772                .unwrap()
773                .as_nanos()
774        ))
775    }
776
777    fn span_end_durations(path: &Path) -> Vec<(String, u64)> {
778        let file = std::fs::File::open(path).unwrap();
779        let reader = BufReader::new(file);
780        reader
781            .lines()
782            .map(|line| line.unwrap())
783            .filter_map(|line| {
784                let value: Value = serde_json::from_str(&line).unwrap();
785                if value.get("kind")?.as_str()? != "span_end" {
786                    return None;
787                }
788                Some((
789                    value["id"].as_str().unwrap().to_string(),
790                    value["duration_ns"].as_u64().unwrap(),
791                ))
792            })
793            .collect()
794    }
795
796    fn read_events(path: &Path) -> Vec<TraceEvent> {
797        let file = std::fs::File::open(path).unwrap();
798        let reader = BufReader::new(file);
799        reader
800            .lines()
801            .map(|line| {
802                let line = line.unwrap();
803                serde_json::from_str(&line).unwrap_or_else(|err| {
804                    panic!("invalid JSONL line `{line}`: {err}");
805                })
806            })
807            .collect()
808    }
809
810    #[test]
811    fn nested_spans_emit_parent_hierarchy_and_durations() {
812        let path = temp_trace("nested");
813        let session =
814            TraceSession::open(&path, ProfileRun::training("model::forward", 1, "cpu")).unwrap();
815
816        let inner_id = {
817            let _outer = session.begin_measurement("Model::forward");
818            std::thread::sleep(std::time::Duration::from_micros(50));
819            let inner = session.begin_span("matmul", SpanKind::Op);
820            std::thread::sleep(std::time::Duration::from_micros(50));
821            inner.id
822        };
823
824        session
825            .record_op(
826                inner_id,
827                OpRecord {
828                    op_name: "matmul",
829                    inputs: &["a".into(), "b".into()],
830                    output: Some("c"),
831                    shape: &[8, 8],
832                    dtype: "f32",
833                    device: "cpu",
834                    duration_ns: 1200,
835                    timestamp_ns: 0,
836                    output_dense_bytes: None,
837                    input_dense_bytes: 0,
838                },
839            )
840            .unwrap();
841
842        session.finish().unwrap();
843
844        let events = read_events(&path);
845        assert!(matches!(events.first(), Some(TraceEvent::Meta { .. })));
846
847        let starts: Vec<&SpanStartEvent> = events
848            .iter()
849            .filter_map(|event| match event {
850                TraceEvent::SpanStart(start) => Some(start),
851                _ => None,
852            })
853            .collect();
854        assert_eq!(starts.len(), 3);
855        assert_eq!(starts[0].name, "model::forward");
856        assert!(starts[0].parent_id.is_none());
857        assert_eq!(starts[1].name, "Model::forward");
858        assert_eq!(starts[1].parent_id.as_deref(), Some("s1"));
859        assert_eq!(starts[2].name, "matmul");
860        assert_eq!(starts[2].parent_id.as_deref(), Some("s2"));
861
862        let ends = span_end_durations(&path);
863        assert_eq!(ends.len(), 3);
864        assert!(ends[0].1 > 0);
865
866        let doc = parse_trace(&path).unwrap();
867        assert_eq!(doc.run.entrypoint, "model::forward");
868        assert_eq!(doc.ops.len(), 1);
869        assert_eq!(doc.ops[0].output_dense_bytes, Some(8 * 8 * 4));
870        assert!(
871            doc.memory.is_empty(),
872            "op metadata must not fabricate tensor lifetime"
873        );
874    }
875
876    #[test]
877    fn measured_region_sync_does_not_overstate_nested_span_timing() {
878        let run = ProfileRun::training("train::update", 2, "cuda:0")
879            .measured_region_device_synchronized();
880
881        assert!(run.measured_region_device_synchronized);
882        assert_eq!(run.timing_mode, TimingMode::Host);
883    }
884
885    #[test]
886    fn record_gradient_round_trips_through_trace_parser() {
887        let path = temp_trace("gradient");
888        let session =
889            TraceSession::open(&path, ProfileRun::training("train::loss", 1, "cpu")).unwrap();
890        session
891            .record_gradient("vb", "encoder.weight", GradientState::Present, Some(0.42))
892            .unwrap();
893        session.finish().unwrap();
894
895        let doc = parse_trace(&path).unwrap();
896        assert_eq!(doc.gradients.len(), 1);
897        assert_eq!(doc.gradients[0].key, "encoder.weight");
898    }
899
900    #[cfg(feature = "candle")]
901    #[test]
902    fn tensor_stats_reduce_finite_non_finite_and_empty_tensors() {
903        use candle_core::{Device, Tensor};
904
905        let path = temp_trace("tensor-stats");
906        let session =
907            TraceSession::open(&path, ProfileRun::training("train::loss", 1, "cpu")).unwrap();
908        let span = session.begin_span("forward", SpanKind::Function);
909        let span_id = span.id();
910        let finite = Tensor::from_vec(vec![1.0f32, -2.0, 3.0, -4.0], (4,), &Device::Cpu).unwrap();
911        session
912            .record_tensor_stats(span_id, "finite", &finite)
913            .unwrap();
914        let corrupt =
915            Tensor::from_vec(vec![1.0f32, f32::NAN, f32::INFINITY], (3,), &Device::Cpu).unwrap();
916        session
917            .record_tensor_stats(span_id, "corrupt", &corrupt)
918            .unwrap();
919        let empty = Tensor::zeros((0,), candle_core::DType::F32, &Device::Cpu).unwrap();
920        session
921            .record_tensor_stats(span_id, "empty", &empty)
922            .unwrap();
923        drop(span);
924        session.finish().unwrap();
925
926        let doc = parse_trace(&path).unwrap();
927        assert_eq!(doc.tensor_stats.len(), 3);
928        let finite = &doc.tensor_stats[0];
929        assert_eq!(finite.elements, 4);
930        assert_eq!(finite.non_finite, 0);
931        assert!((finite.rms - 7.5f64.sqrt()).abs() < 1e-6);
932        assert_eq!(finite.abs_max, 4.0);
933        assert_eq!(finite.mean, -0.5);
934        assert_eq!(doc.tensor_stats[1].non_finite, 2);
935        assert_eq!(doc.tensor_stats[2].elements, 0);
936        assert_eq!(doc.tensor_stats[2].non_finite, 0);
937        assert_eq!(doc.tensor_stats[2].rms, 0.0);
938    }
939
940    #[test]
941    fn scalar_observations_round_trip_as_single_element_stats() {
942        let path = temp_trace("scalar");
943        let session =
944            TraceSession::open(&path, ProfileRun::training("train::update", 1, "cpu")).unwrap();
945        let span = session.begin_span("loss", SpanKind::Function);
946        let span_id = span.id();
947        session.record_scalar(span_id, "loss/total", 1.25).unwrap();
948        session
949            .record_scalar(span_id, "loss/aux", f64::NAN)
950            .unwrap();
951        assert!(session.record_scalar(span_id, "  ", 1.0).is_err());
952        drop(span);
953        session.finish().unwrap();
954
955        let doc = parse_trace(&path).unwrap();
956        assert_eq!(doc.tensor_stats.len(), 2);
957        let total = &doc.tensor_stats[0];
958        assert_eq!(total.label, "loss/total");
959        assert_eq!(total.elements, 1);
960        assert_eq!(total.non_finite, 0);
961        assert_eq!(total.mean, 1.25);
962        assert_eq!(total.rms, 1.25);
963        assert!(total.shape.is_empty());
964        let aux = &doc.tensor_stats[1];
965        assert_eq!(aux.non_finite, 1);
966        assert_eq!(aux.mean, 0.0);
967    }
968
969    #[test]
970    fn completed_host_span_is_root_attached_stack_neutral_and_overlap_valid() {
971        let path = temp_trace("completed-host-span");
972        let session =
973            TraceSession::open(&path, ProfileRun::inference("serve::request", 1, "cpu")).unwrap();
974
975        let measured = session.begin_measurement("request");
976        let live = session.begin_span("main-thread", SpanKind::Function);
977        let stack_before = session.inner.borrow().span_stack.clone();
978        let start_ns = session.elapsed_ns();
979        let end_ns = loop {
980            let elapsed = session.elapsed_ns();
981            if elapsed > start_ns {
982                break elapsed;
983            }
984            std::hint::spin_loop();
985        };
986        let duration_ns = end_ns - start_ns;
987
988        let completed_id = session
989            .record_completed_host_span(
990                "background-work",
991                SpanKind::Function,
992                start_ns,
993                duration_ns,
994            )
995            .unwrap();
996        assert_eq!(session.inner.borrow().span_stack, stack_before);
997
998        drop(live);
999        drop(measured);
1000        session.finish().unwrap();
1001
1002        let doc = parse_trace(&path).unwrap();
1003        let completed = doc
1004            .spans
1005            .iter()
1006            .find(|span| span.id == span_id_string(completed_id.0))
1007            .unwrap();
1008        let measured = doc.spans.iter().find(|span| span.measured).unwrap();
1009        assert_eq!(completed.parent_id.as_deref(), Some("s1"));
1010        assert_eq!(completed.start_ns, start_ns);
1011        assert_eq!(completed.duration_ns, duration_ns);
1012        assert!(completed.closed);
1013        assert!(completed.start_ns < measured.start_ns.saturating_add(measured.duration_ns));
1014        assert!(measured.start_ns < completed.start_ns.saturating_add(completed.duration_ns));
1015
1016        let health = analyze_health(&doc);
1017        assert!(health.structurally_valid, "{:?}", health.issues);
1018        assert!(health.capture_complete);
1019    }
1020
1021    #[test]
1022    fn completed_host_span_requires_positive_duration() {
1023        let path = temp_trace("completed-host-span-zero");
1024        let session =
1025            TraceSession::open(&path, ProfileRun::inference("serve::request", 1, "cpu")).unwrap();
1026        let error = session
1027            .record_completed_host_span("empty", SpanKind::Function, 0, 0)
1028            .unwrap_err();
1029        assert!(error.to_string().contains("duration must be positive"));
1030
1031        let overflow = session
1032            .record_completed_host_span("overflow", SpanKind::Function, u64::MAX, 1)
1033            .unwrap_err();
1034        assert!(overflow.to_string().contains("overflows"));
1035
1036        let future = session
1037            .record_completed_host_span(
1038                "future",
1039                SpanKind::Function,
1040                session.elapsed_ns().saturating_add(60_000_000_000),
1041                1,
1042            )
1043            .unwrap_err();
1044        assert!(future.to_string().contains("ends in the future"));
1045
1046        let before_session = session
1047            .inner
1048            .borrow()
1049            .probe_started
1050            .checked_sub(std::time::Duration::from_nanos(1))
1051            .unwrap();
1052        assert!(session.host_timestamp_ns(before_session).is_err());
1053        let now = Instant::now();
1054        assert!(session.host_timestamp_ns(now).unwrap() <= session.elapsed_ns());
1055        session.finish().unwrap();
1056    }
1057
1058    #[test]
1059    fn sticky_error_poisons_event_writes_and_finish() {
1060        let path = temp_trace("sticky-error");
1061        let session =
1062            TraceSession::open(&path, ProfileRun::training("train::update", 1, "cpu")).unwrap();
1063        session.inner.borrow_mut().sticky_error = Some("boom".into());
1064
1065        let error = session
1066            .record_op(
1067                SpanId(1),
1068                OpRecord {
1069                    op_name: "matmul",
1070                    inputs: &[],
1071                    output: None,
1072                    shape: &[2, 2],
1073                    dtype: "f32",
1074                    device: "cpu",
1075                    duration_ns: 100,
1076                    timestamp_ns: 0,
1077                    output_dense_bytes: None,
1078                    input_dense_bytes: 0,
1079                },
1080            )
1081            .unwrap_err();
1082        assert!(error.to_string().contains("previously failed"));
1083
1084        let error = session.finish().unwrap_err();
1085        assert!(error.to_string().contains("previously failed"));
1086    }
1087}