Skip to main content

dynamo_data_gen/request_trace/
mooncake.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use crate::{MooncakeRow, RollingHashIdMapper};
5use anyhow::{Context, Result, anyhow, bail};
6
7use super::load::RequestEntry;
8
9/// Streams each request through a Mooncake-compatible row into the replay builder.
10///
11/// This is an in-memory compatibility layer; it does not write a Mooncake trace.
12pub fn lower_mooncake_rows<F>(mut requests: Vec<RequestEntry>, mut emit: F) -> Result<usize>
13where
14    F: FnMut(usize, MooncakeRow) -> Result<()>,
15{
16    let global_start_ms = requests
17        .iter()
18        .map(|request| request.start_ms)
19        .min()
20        .ok_or_else(|| anyhow!("no request records to convert"))?;
21    let trace_block_size = requests[0].replay.trace_block_size;
22    for request in &requests {
23        if request.replay.trace_block_size != trace_block_size {
24            bail!(
25                "mixed replay trace_block_size values are not supported: {} and {}",
26                trace_block_size,
27                request.replay.trace_block_size
28            );
29        }
30    }
31
32    requests.sort_by(|left, right| {
33        (left.start_ms, left.end_ms, &left.request.request_id).cmp(&(
34            right.start_ms,
35            right.end_ms,
36            &right.request.request_id,
37        ))
38    });
39
40    let mut mapper = RollingHashIdMapper::new(trace_block_size);
41    for request in requests {
42        let hash_ids = mapper.ids_for_sequence_hashes(&request.replay.input_sequence_hashes);
43        let output_length = request.request.output_tokens.ok_or_else(|| {
44            anyhow!(
45                "request {} is missing output length",
46                request.request.request_id
47            )
48        })?;
49        emit(
50            trace_block_size,
51            MooncakeRow {
52                session_id: None,
53                input_length: Some(request.replay.input_length),
54                output_length: Some(
55                    usize::try_from(output_length)
56                        .context("output length does not fit in usize")?,
57                ),
58                hash_ids: Some(hash_ids),
59                timestamp: Some((request.start_ms - global_start_ms) as f64),
60                delay: None,
61                ..Default::default()
62            },
63        )?;
64    }
65
66    Ok(trace_block_size)
67}
68
69#[cfg(test)]
70mod tests {
71    use super::*;
72    use crate::request_trace::load::{
73        RequestEntry, RequestTraceReplayMetrics, RequestTraceRequestMetrics,
74    };
75
76    fn request(
77        request_id: &str,
78        start_ms: i64,
79        end_ms: i64,
80        sequence_hashes: Vec<u64>,
81    ) -> RequestEntry {
82        RequestEntry {
83            start_ms,
84            end_ms,
85            agent_context: None,
86            request: RequestTraceRequestMetrics {
87                request_id: request_id.to_string(),
88                output_tokens: Some(5),
89                request_received_ms: Some(start_ms as u64),
90                total_time_ms: Some((end_ms - start_ms) as f64),
91                ..Default::default()
92            },
93            replay: RequestTraceReplayMetrics {
94                trace_block_size: 2,
95                input_length: sequence_hashes.len() * 2,
96                input_sequence_hashes: sequence_hashes,
97            },
98        }
99    }
100
101    #[test]
102    fn lowering_preserves_timestamp_offsets_and_parallel_requests() {
103        let requests = vec![
104            request("req-a", 1_000, 1_100, vec![11, 22]),
105            request("req-b", 1_000, 1_700, vec![22]),
106            request("req-c", 1_500, 1_600, vec![11, 33]),
107        ];
108
109        let mut entries = Vec::new();
110        lower_mooncake_rows(requests, |_, row| {
111            entries.push(row);
112            Ok(())
113        })
114        .unwrap();
115
116        assert_eq!(entries.len(), 3);
117        assert_eq!(entries[0].timestamp, Some(0.0));
118        assert_eq!(entries[1].timestamp, Some(0.0));
119        assert_eq!(entries[2].timestamp, Some(500.0));
120        assert!(entries.iter().all(|entry| entry.delay.is_none()));
121        assert!(entries.iter().all(|entry| entry.session_id.is_none()));
122        assert_eq!(
123            entries[0].hash_ids.as_ref().unwrap()[0],
124            entries[2].hash_ids.as_ref().unwrap()[0]
125        );
126    }
127}