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