dynamo_data_gen/request_trace/
mooncake.rs1use crate::{MooncakeRow, RollingHashIdMapper};
5use anyhow::{Context, Result, anyhow, bail};
6
7use super::load::RequestEntry;
8
9pub 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}