Skip to main content

dynamo_mocker/loadgen/
driver.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::cmp::Ordering;
5use std::collections::BinaryHeap;
6
7use anyhow::{Context, Result, anyhow, bail};
8use rand::SeedableRng;
9use rand::rngs::StdRng;
10use rustc_hash::FxHashMap;
11use uuid::Uuid;
12
13use super::trace::validate_synthesizable_prompt;
14use super::types::{
15    AgenticTrace, CompactReadyTurn, ReadyTurn, ReplayRequestHashes, ReplayRequestPayload, Trace,
16};
17use super::{SYNTHETIC_OUTPUT_SEED, planned_output_token_ids};
18use crate::common::protocols::DirectRequest;
19
20#[derive(Debug)]
21enum SchedulingPolicy {
22    Trace,
23    Concurrency(ConcurrencyState),
24    Agentic(AgenticState),
25}
26
27#[derive(Debug)]
28struct ConcurrencyState {
29    max_active_sessions: usize,
30    next_pending_session: usize,
31    active_sessions: usize,
32}
33
34#[derive(Debug)]
35struct AgenticState {
36    remaining_dependencies: Vec<usize>,
37    ready_after_ms: Vec<f64>,
38    dependents: FxHashMap<String, Vec<usize>>,
39}
40
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42enum PromptMode {
43    Full,
44    DeltaCumulative,
45}
46
47#[derive(Debug, Clone, Copy, PartialEq, Eq)]
48enum TurnOutcome {
49    Completed,
50    Rejected,
51    Cancelled,
52}
53
54#[derive(Debug)]
55struct TurnResolution {
56    request_id: Option<String>,
57    session_ended: bool,
58}
59
60#[derive(Debug)]
61struct SessionRuntime {
62    session_id: String,
63    turns: Vec<TurnRuntime>,
64    cumulative_tokens: Vec<u32>,
65    next_turn_index: usize,
66    next_ready_at_ms: Option<f64>,
67    in_flight: Option<Uuid>,
68}
69
70#[derive(Debug)]
71enum PromptTokens {
72    // Full-prompt traces stay in their compact on-disk representation until
73    // dispatch. Delta-cumulative traces remain eager because later turns append
74    // generated output to already-materialized session history.
75    Deferred {
76        input_length: usize,
77        hash_ids: Vec<u32>,
78    },
79    Materialized(Vec<u32>),
80}
81
82impl PromptTokens {
83    fn deferred(input_length: usize, hash_ids: Vec<u32>, trace_block_size: usize) -> Result<Self> {
84        validate_synthesizable_prompt(input_length, &hash_ids, trace_block_size)?;
85        Ok(Self::Deferred {
86            input_length,
87            hash_ids,
88        })
89    }
90
91    fn input_length(&self) -> usize {
92        match self {
93            Self::Deferred { input_length, .. } => *input_length,
94            Self::Materialized(tokens) => tokens.len(),
95        }
96    }
97
98    fn take_deferred(&mut self) -> (usize, Vec<u32>) {
99        match self {
100            Self::Deferred {
101                input_length,
102                hash_ids,
103            } => (*input_length, std::mem::take(hash_ids)),
104            Self::Materialized(_) => {
105                unreachable!("full-prompt turns must retain their deferred representation")
106            }
107        }
108    }
109
110    fn materialized(&self) -> &[u32] {
111        match self {
112            Self::Deferred { .. } => {
113                unreachable!("delta-cumulative prompts are materialized during driver setup")
114            }
115            Self::Materialized(tokens) => tokens,
116        }
117    }
118}
119
120#[derive(Debug)]
121struct TurnRuntime {
122    request_id: Option<String>,
123    replay_key: Option<String>,
124    prompt_tokens: PromptTokens,
125    max_output_tokens: usize,
126    output_token_ids: Option<Vec<u32>>,
127    delay_after_previous_ms: f64,
128    priority: i32,
129    strict_priority: u32,
130    policy_class: Option<String>,
131    #[cfg(any(test, feature = "replay-bench"))]
132    deterministic_request_id: Option<Uuid>,
133}
134
135#[derive(Debug, Clone, Copy)]
136struct InFlightTurn {
137    session_index: usize,
138    turn_index: usize,
139    emitted_output_tokens: usize,
140}
141
142#[derive(Debug, Clone, Copy)]
143struct ReadySession {
144    ready_at_ms: f64,
145    session_index: usize,
146    turn_index: usize,
147}
148
149impl PartialEq for ReadySession {
150    fn eq(&self, other: &Self) -> bool {
151        self.ready_at_ms.to_bits() == other.ready_at_ms.to_bits()
152            && self.session_index == other.session_index
153            && self.turn_index == other.turn_index
154    }
155}
156
157impl Eq for ReadySession {}
158
159impl Ord for ReadySession {
160    fn cmp(&self, other: &Self) -> Ordering {
161        other
162            .ready_at_ms
163            .total_cmp(&self.ready_at_ms)
164            .then_with(|| other.session_index.cmp(&self.session_index))
165            .then_with(|| other.turn_index.cmp(&self.turn_index))
166    }
167}
168
169impl PartialOrd for ReadySession {
170    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
171        Some(self.cmp(other))
172    }
173}
174
175impl SchedulingPolicy {
176    fn schedules_sequential_turns(&self) -> bool {
177        !matches!(self, Self::Agentic(_))
178    }
179
180    fn arrival_timestamp_ms(&self, scheduled_ready_at_ms: f64) -> Option<f64> {
181        match self {
182            Self::Concurrency(_) => None,
183            Self::Trace | Self::Agentic(_) => Some(scheduled_ready_at_ms),
184        }
185    }
186
187    fn dispatch_limit(&self, requested: usize, in_flight: usize) -> usize {
188        match self {
189            Self::Concurrency(state) => {
190                requested.min(state.max_active_sessions.saturating_sub(in_flight))
191            }
192            Self::Trace | Self::Agentic(_) => requested,
193        }
194    }
195
196    fn at_dispatch_capacity(&self, in_flight: usize) -> bool {
197        matches!(
198            self,
199            Self::Concurrency(state) if in_flight >= state.max_active_sessions
200        )
201    }
202}
203
204impl ConcurrencyState {
205    fn new(max_active_sessions: usize) -> Self {
206        Self {
207            max_active_sessions,
208            next_pending_session: 0,
209            active_sessions: 0,
210        }
211    }
212
213    fn activate_pending(
214        &mut self,
215        sessions: &mut [SessionRuntime],
216        ready_sessions: &mut BinaryHeap<ReadySession>,
217        now_ms: f64,
218    ) {
219        while self.active_sessions < self.max_active_sessions
220            && self.next_pending_session < sessions.len()
221        {
222            let session_index = self.next_pending_session;
223            self.next_pending_session += 1;
224            let session = &mut sessions[session_index];
225            let turn_index = session.next_turn_index;
226            session.next_ready_at_ms = Some(now_ms);
227            ready_sessions.push(ReadySession {
228                ready_at_ms: now_ms,
229                session_index,
230                turn_index,
231            });
232            self.active_sessions += 1;
233        }
234    }
235
236    fn on_session_finished(
237        &mut self,
238        sessions: &mut [SessionRuntime],
239        ready_sessions: &mut BinaryHeap<ReadySession>,
240        now_ms: f64,
241    ) {
242        self.active_sessions = self.active_sessions.saturating_sub(1);
243        self.activate_pending(sessions, ready_sessions, now_ms);
244    }
245}
246
247impl AgenticState {
248    fn release_dependents(
249        &mut self,
250        sessions: &mut [SessionRuntime],
251        ready_sessions: &mut BinaryHeap<ReadySession>,
252        request_id: &str,
253        now_ms: f64,
254    ) {
255        let Some(dependent_sessions) = self.dependents.get(request_id).cloned() else {
256            return;
257        };
258        for session_index in dependent_sessions {
259            let Some(remaining) = self.remaining_dependencies.get_mut(session_index) else {
260                continue;
261            };
262            if *remaining == 0 {
263                continue;
264            }
265            *remaining -= 1;
266            if let Some(ready_after_ms) = self.ready_after_ms.get_mut(session_index) {
267                *ready_after_ms = ready_after_ms.max(now_ms);
268            }
269            if *remaining != 0 {
270                continue;
271            }
272
273            let Some(session) = sessions.get_mut(session_index) else {
274                continue;
275            };
276            if session.in_flight.is_some()
277                || session.next_turn_index >= session.turns.len()
278                || session.next_ready_at_ms.is_some()
279            {
280                continue;
281            }
282            let turn_index = session.next_turn_index;
283            let ready_at_ms = self.ready_after_ms[session_index]
284                + session.turns[turn_index].delay_after_previous_ms;
285            session.next_ready_at_ms = Some(ready_at_ms);
286            ready_sessions.push(ReadySession {
287                ready_at_ms,
288                session_index,
289                turn_index,
290            });
291        }
292    }
293}
294
295#[derive(Debug)]
296pub struct WorkloadDriver {
297    policy: SchedulingPolicy,
298    prompt_mode: PromptMode,
299    emit_session_metadata: bool,
300    trace_block_size: usize,
301    engine_block_size: u32,
302    include_replay_hashes: bool,
303    sessions: Vec<SessionRuntime>,
304    in_flight: FxHashMap<Uuid, InFlightTurn>,
305    ready_sessions: BinaryHeap<ReadySession>,
306}
307
308impl WorkloadDriver {
309    pub(crate) fn new_trace(trace: Trace, engine_block_size: usize) -> Result<Self> {
310        Self::new(
311            trace,
312            engine_block_size,
313            SchedulingPolicy::Trace,
314            PromptMode::Full,
315            true,
316        )
317    }
318
319    pub(crate) fn new_trace_without_replay_hashes(
320        trace: Trace,
321        engine_block_size: usize,
322        accumulate_session_deltas: bool,
323    ) -> Result<Self> {
324        trace.validate_for_trace_mode()?;
325        let prompt_mode = if accumulate_session_deltas {
326            PromptMode::DeltaCumulative
327        } else {
328            PromptMode::Full
329        };
330        Self::new(
331            trace,
332            engine_block_size,
333            SchedulingPolicy::Trace,
334            prompt_mode,
335            false,
336        )
337    }
338
339    pub(crate) fn new_trace_accumulating_deltas(
340        trace: Trace,
341        engine_block_size: usize,
342    ) -> Result<Self> {
343        Self::new(
344            trace,
345            engine_block_size,
346            SchedulingPolicy::Trace,
347            PromptMode::DeltaCumulative,
348            true,
349        )
350    }
351
352    /// Build a closed-loop concurrency driver. `max_in_flight` is the *session* cap
353    /// (depth-first): a session holds its slot across all turns + think-time, and new
354    /// sessions are admitted only while fewer than `max_in_flight` are active.
355    pub(crate) fn new_concurrency(
356        trace: Trace,
357        engine_block_size: usize,
358        max_in_flight: usize,
359    ) -> Result<Self> {
360        Self::new(
361            trace,
362            engine_block_size,
363            SchedulingPolicy::Concurrency(ConcurrencyState::new(max_in_flight)),
364            PromptMode::Full,
365            true,
366        )
367    }
368
369    pub(crate) fn new_concurrency_without_replay_hashes(
370        trace: Trace,
371        engine_block_size: usize,
372        max_in_flight: usize,
373        accumulate_session_deltas: bool,
374    ) -> Result<Self> {
375        trace.validate_for_concurrency_mode()?;
376        let prompt_mode = if accumulate_session_deltas {
377            PromptMode::DeltaCumulative
378        } else {
379            PromptMode::Full
380        };
381        Self::new(
382            trace,
383            engine_block_size,
384            SchedulingPolicy::Concurrency(ConcurrencyState::new(max_in_flight)),
385            prompt_mode,
386            false,
387        )
388    }
389
390    pub(crate) fn new_concurrency_accumulating_deltas(
391        trace: Trace,
392        engine_block_size: usize,
393        max_in_flight: usize,
394    ) -> Result<Self> {
395        Self::new(
396            trace,
397            engine_block_size,
398            SchedulingPolicy::Concurrency(ConcurrencyState::new(max_in_flight)),
399            PromptMode::DeltaCumulative,
400            true,
401        )
402    }
403
404    pub(crate) fn new_agentic_trace(trace: AgenticTrace, engine_block_size: usize) -> Result<Self> {
405        Self::new_agentic_trace_with_replay_hashes(trace, engine_block_size, true)
406    }
407
408    pub(crate) fn new_agentic_trace_without_replay_hashes(
409        trace: AgenticTrace,
410        engine_block_size: usize,
411    ) -> Result<Self> {
412        Self::new_agentic_trace_with_replay_hashes(trace, engine_block_size, false)
413    }
414
415    fn new_agentic_trace_with_replay_hashes(
416        trace: AgenticTrace,
417        engine_block_size: usize,
418        include_replay_hashes: bool,
419    ) -> Result<Self> {
420        if engine_block_size == 0 {
421            bail!("engine_block_size must be greater than 0");
422        }
423        let engine_block_size_u32 =
424            u32::try_from(engine_block_size).context("engine_block_size does not fit in u32")?;
425        let trace_block_size = trace.block_size;
426
427        let mut dependents: FxHashMap<String, Vec<usize>> = FxHashMap::default();
428        let mut remaining_dependencies = Vec::with_capacity(trace.turns.len());
429        let mut ready_after_ms = Vec::with_capacity(trace.turns.len());
430        let mut sessions = Vec::with_capacity(trace.turns.len());
431        let mut output_rng = StdRng::seed_from_u64(SYNTHETIC_OUTPUT_SEED);
432
433        for (session_index, mut turn) in trace.turns.into_iter().enumerate() {
434            for dependency in &turn.wait_for {
435                dependents
436                    .entry(dependency.clone())
437                    .or_default()
438                    .push(session_index);
439            }
440            remaining_dependencies.push(turn.wait_for.len());
441            ready_after_ms.push(0.0);
442
443            let prompt_tokens = PromptTokens::deferred(
444                turn.input_length,
445                std::mem::take(&mut turn.hash_ids),
446                trace_block_size,
447            )?;
448            let output_token_ids = Some(planned_output_token_ids(
449                turn.output_token_ids,
450                turn.max_output_tokens,
451                &mut output_rng,
452            ));
453            let next_ready_at_ms = if turn.wait_for.is_empty() {
454                Some(turn.first_ready_timestamp_ms.unwrap_or(0.0))
455            } else {
456                None
457            };
458            sessions.push(SessionRuntime {
459                session_id: turn.session_id,
460                turns: vec![TurnRuntime {
461                    request_id: Some(turn.request_id),
462                    replay_key: turn.replay_key,
463                    prompt_tokens,
464                    max_output_tokens: turn.max_output_tokens,
465                    output_token_ids,
466                    delay_after_previous_ms: turn.delay_after_dependencies_ms,
467                    priority: turn.priority,
468                    strict_priority: turn.strict_priority,
469                    policy_class: turn.policy_class,
470                    #[cfg(feature = "replay-bench")]
471                    deterministic_request_id: Some(Uuid::from_u128(session_index as u128 + 1)),
472                    #[cfg(all(test, not(feature = "replay-bench")))]
473                    deterministic_request_id: None,
474                }],
475                cumulative_tokens: Vec::new(),
476                next_turn_index: 0,
477                next_ready_at_ms,
478                in_flight: None,
479            });
480        }
481
482        let ready_sessions = sessions
483            .iter()
484            .enumerate()
485            .filter_map(|(session_index, session)| {
486                Some(ReadySession {
487                    ready_at_ms: session.next_ready_at_ms?,
488                    session_index,
489                    turn_index: session.next_turn_index,
490                })
491            })
492            .collect();
493
494        Ok(Self {
495            policy: SchedulingPolicy::Agentic(AgenticState {
496                remaining_dependencies,
497                ready_after_ms,
498                dependents,
499            }),
500            prompt_mode: PromptMode::Full,
501            emit_session_metadata: true,
502            trace_block_size,
503            engine_block_size: engine_block_size_u32,
504            include_replay_hashes,
505            sessions,
506            in_flight: FxHashMap::default(),
507            ready_sessions,
508        })
509    }
510
511    fn new(
512        trace: Trace,
513        engine_block_size: usize,
514        policy: SchedulingPolicy,
515        prompt_mode: PromptMode,
516        include_replay_hashes: bool,
517    ) -> Result<Self> {
518        if engine_block_size == 0 {
519            bail!("engine_block_size must be greater than 0");
520        }
521        let engine_block_size_u32 =
522            u32::try_from(engine_block_size).context("engine_block_size does not fit in u32")?;
523        let trace_block_size = trace.block_size;
524        let is_concurrency = matches!(&policy, SchedulingPolicy::Concurrency(_));
525        let mut output_rng = StdRng::seed_from_u64(SYNTHETIC_OUTPUT_SEED);
526        #[cfg(feature = "replay-bench")]
527        let mut next_deterministic_request_id = 1_u128;
528        let sessions: Vec<SessionRuntime> = trace
529            .sessions
530            .into_iter()
531            .map(|session| -> Result<SessionRuntime> {
532                let next_ready_at_ms = if is_concurrency {
533                    None
534                } else {
535                    Some(session.first_arrival_timestamp_ms.unwrap_or(0.0))
536                };
537                let turns = session
538                    .turns
539                    .into_iter()
540                    .map(|mut turn| -> Result<TurnRuntime> {
541                        let prompt_tokens = match prompt_mode {
542                            PromptMode::Full => PromptTokens::deferred(
543                                turn.input_length,
544                                std::mem::take(&mut turn.hash_ids),
545                                trace_block_size,
546                            )?,
547                            PromptMode::DeltaCumulative => PromptTokens::Materialized(
548                                turn.synthesize_tokens(trace_block_size)?,
549                            ),
550                        };
551                        let output_token_ids = Some(planned_output_token_ids(
552                            turn.output_token_ids,
553                            turn.max_output_tokens,
554                            &mut output_rng,
555                        ));
556                        #[cfg(feature = "replay-bench")]
557                        let deterministic_request_id = {
558                            let request_id = Uuid::from_u128(next_deterministic_request_id);
559                            next_deterministic_request_id = next_deterministic_request_id
560                                .checked_add(1)
561                                .expect("deterministic replay request UUID overflow");
562                            Some(request_id)
563                        };
564                        Ok(TurnRuntime {
565                            request_id: None,
566                            prompt_tokens,
567                            replay_key: turn.replay_key,
568                            max_output_tokens: turn.max_output_tokens,
569                            output_token_ids,
570                            delay_after_previous_ms: turn.delay_after_previous_ms,
571                            priority: turn.priority,
572                            strict_priority: turn.strict_priority,
573                            policy_class: turn.policy_class,
574                            #[cfg(feature = "replay-bench")]
575                            deterministic_request_id,
576                            #[cfg(all(test, not(feature = "replay-bench")))]
577                            deterministic_request_id: None,
578                        })
579                    })
580                    .collect::<Result<Vec<_>>>()?;
581                let cumulative_capacity = if prompt_mode == PromptMode::DeltaCumulative {
582                    turns
583                        .iter()
584                        .map(|turn| {
585                            turn.prompt_tokens.input_length()
586                                + turn
587                                    .output_token_ids
588                                    .as_ref()
589                                    .map_or(0, |output| output.len())
590                        })
591                        .sum()
592                } else {
593                    0
594                };
595                Ok(SessionRuntime {
596                    session_id: session.session_id,
597                    turns,
598                    cumulative_tokens: Vec::with_capacity(cumulative_capacity),
599                    next_turn_index: 0,
600                    next_ready_at_ms,
601                    in_flight: None,
602                })
603            })
604            .collect::<Result<Vec<_>>>()?;
605
606        let ready_sessions = sessions
607            .iter()
608            .enumerate()
609            .filter_map(|(session_index, session)| {
610                Some(ReadySession {
611                    ready_at_ms: session.next_ready_at_ms?,
612                    session_index,
613                    turn_index: session.next_turn_index,
614                })
615            })
616            .collect();
617
618        let mut driver = Self {
619            policy,
620            prompt_mode,
621            emit_session_metadata: true,
622            trace_block_size,
623            engine_block_size: engine_block_size_u32,
624            include_replay_hashes,
625            sessions,
626            in_flight: FxHashMap::default(),
627            ready_sessions,
628        };
629        if let SchedulingPolicy::Concurrency(state) = &mut driver.policy {
630            state.activate_pending(&mut driver.sessions, &mut driver.ready_sessions, 0.0);
631        }
632        Ok(driver)
633    }
634
635    /// Use stable monotonically increasing UUIDs for replay parity fixtures.
636    /// This is unavailable in production builds so normal request identity and
637    /// randomness remain unchanged.
638    #[cfg(any(test, feature = "replay-bench"))]
639    pub fn with_deterministic_request_ids(mut self, first_id: u128) -> Self {
640        let mut next_id = first_id;
641        for session in &mut self.sessions {
642            for turn in &mut session.turns {
643                turn.deterministic_request_id = Some(Uuid::from_u128(next_id));
644                next_id = next_id
645                    .checked_add(1)
646                    .expect("deterministic replay request UUID overflow");
647            }
648        }
649        self
650    }
651
652    fn request_uuid(&self, _session_index: usize, _turn_index: usize) -> Uuid {
653        #[cfg(any(test, feature = "replay-bench"))]
654        if let Some(request_id) =
655            self.sessions[_session_index].turns[_turn_index].deterministic_request_id
656        {
657            return request_id;
658        }
659
660        Uuid::new_v4()
661    }
662
663    pub(crate) fn without_session_metadata(mut self) -> Self {
664        self.emit_session_metadata = false;
665        self
666    }
667
668    /// Failure-path companion: release a cap slot and terminate the owning session.
669    /// No-op if `on_complete` already ran. Used when a request task is cancelled
670    /// or panics before reaching `on_complete`.
671    ///
672    /// Terminating the session (marking it exhausted) prevents `run_workload` from
673    /// deadlocking: `pop_ready` skips sessions with `in_flight.is_some()`, so a
674    /// leaked session would leave `is_drained` stuck at `false` forever.
675    pub fn release_cap_slot(&mut self, request_uuid: Uuid, now_ms: f64) {
676        let Ok(Some(resolution)) = self.resolve_turn(request_uuid, now_ms, TurnOutcome::Cancelled)
677        else {
678            return;
679        };
680        self.apply_resolution(resolution, now_ms);
681    }
682
683    pub fn pop_ready(&mut self, now_ms: f64, limit: usize) -> Vec<ReadyTurn> {
684        self.pop_ready_compact(now_ms, limit)
685            .into_iter()
686            .map(CompactReadyTurn::into_ready_turn)
687            .collect()
688    }
689
690    pub(crate) fn pop_ready_compact(&mut self, now_ms: f64, limit: usize) -> Vec<CompactReadyTurn> {
691        let effective_limit = self.policy.dispatch_limit(limit, self.in_flight.len());
692        if effective_limit == 0 {
693            return Vec::new();
694        }
695
696        let mut emitted = Vec::new();
697        while emitted.len() < effective_limit {
698            let Some(ready_session) = self.ready_sessions.pop() else {
699                break;
700            };
701            if ready_session.ready_at_ms > now_ms {
702                self.ready_sessions.push(ready_session);
703                break;
704            }
705
706            let session_index = ready_session.session_index;
707            let Some((turn_index, scheduled_ready_at_ms)) = self
708                .sessions
709                .get(session_index)
710                .filter(|session| {
711                    session.in_flight.is_none()
712                        && session.next_turn_index == ready_session.turn_index
713                        && session.next_ready_at_ms == Some(ready_session.ready_at_ms)
714                })
715                .map(|session| {
716                    (
717                        session.next_turn_index,
718                        session
719                            .next_ready_at_ms
720                            .expect("ready session must have a timestamp"),
721                    )
722                })
723            else {
724                continue;
725            };
726            let request_uuid = self.request_uuid(session_index, turn_index);
727            let session = &mut self.sessions[session_index];
728            let turn = &mut session.turns[turn_index];
729            let arrival_timestamp_ms = self.policy.arrival_timestamp_ms(scheduled_ready_at_ms);
730            let (request, replay_hashes) = match self.prompt_mode {
731                PromptMode::Full => {
732                    let (input_length, hash_ids) = turn.prompt_tokens.take_deferred();
733                    let request_metadata = DirectRequest {
734                        tokens: Vec::new(),
735                        max_output_tokens: turn.max_output_tokens,
736                        output_token_ids: turn.output_token_ids.take(),
737                        uuid: Some(request_uuid),
738                        dp_rank: 0,
739                        arrival_timestamp_ms,
740                        priority: turn.priority,
741                        strict_priority: turn.strict_priority,
742                        policy_class: turn.policy_class.clone(),
743                    };
744                    let request = ReplayRequestPayload::deferred(
745                        request_metadata,
746                        input_length,
747                        hash_ids,
748                        self.trace_block_size,
749                    );
750                    // The router needs engine-block hashes at arrival, but it
751                    // does not need to retain the expanded prompt. Materialize
752                    // once transiently for hashing, then keep only the compact
753                    // payload until a worker admission.
754                    // TODO: Derive engine-block hashes directly from the compact
755                    // trace blocks so immediate dispatch does not materialize
756                    // the prompt once for routing and again for admission.
757                    // Preserve `ReplayRequestHashes::from_tokens` semantics when
758                    // trace and engine block sizes differ.
759                    let replay_hashes = self.include_replay_hashes.then(|| {
760                        let request_tokens = request.prompt_tokens();
761                        ReplayRequestHashes::from_tokens(&request_tokens, self.engine_block_size)
762                    });
763                    (request, replay_hashes)
764                }
765                PromptMode::DeltaCumulative => {
766                    session
767                        .cumulative_tokens
768                        .extend_from_slice(turn.prompt_tokens.materialized());
769                    let request_tokens = session.cumulative_tokens.clone();
770                    let replay_hashes = self.include_replay_hashes.then(|| {
771                        ReplayRequestHashes::from_tokens(&request_tokens, self.engine_block_size)
772                    });
773                    let request = ReplayRequestPayload::materialized(DirectRequest {
774                        tokens: request_tokens,
775                        max_output_tokens: turn.max_output_tokens,
776                        output_token_ids: turn.output_token_ids.clone(),
777                        uuid: Some(request_uuid),
778                        dp_rank: 0,
779                        arrival_timestamp_ms,
780                        priority: turn.priority,
781                        strict_priority: turn.strict_priority,
782                        policy_class: turn.policy_class.clone(),
783                    });
784                    (request, replay_hashes)
785                }
786            };
787            session.in_flight = Some(request_uuid);
788            session.next_ready_at_ms = None;
789            self.in_flight.insert(
790                request_uuid,
791                InFlightTurn {
792                    session_index,
793                    turn_index,
794                    emitted_output_tokens: 0,
795                },
796            );
797            emitted.push(CompactReadyTurn {
798                request_uuid,
799                session_id: session.session_id.clone(),
800                turn_index,
801                replay_key: turn.replay_key.clone(),
802                scheduled_ready_at_ms,
803                replay_hashes,
804                emit_session_metadata: self.emit_session_metadata,
805                request,
806            });
807        }
808        emitted
809    }
810
811    pub fn on_output_token(&mut self, request_uuid: Uuid, token_id: u32) -> Result<()> {
812        if self.prompt_mode == PromptMode::Full {
813            return Ok(());
814        }
815        let in_flight = self
816            .in_flight
817            .get(&request_uuid)
818            .copied()
819            .ok_or_else(|| anyhow!("unknown workload request output for {request_uuid}"))?;
820
821        let turn = &self.sessions[in_flight.session_index].turns[in_flight.turn_index];
822        let planned_output_tokens = turn
823            .output_token_ids
824            .as_ref()
825            .expect("delta turns must have planned output tokens");
826        let expected_token = planned_output_tokens
827            .get(in_flight.emitted_output_tokens)
828            .ok_or_else(|| {
829                anyhow!(
830                    "workload request {request_uuid} emitted more than {} planned output tokens",
831                    planned_output_tokens.len()
832                )
833            })?;
834        if token_id != *expected_token {
835            bail!(
836                "workload request {request_uuid} emitted token {token_id} at position {}, expected {}",
837                in_flight.emitted_output_tokens,
838                expected_token
839            );
840        }
841
842        let in_flight = self
843            .in_flight
844            .get_mut(&request_uuid)
845            .expect("validated in-flight request must still exist");
846        in_flight.emitted_output_tokens = in_flight
847            .emitted_output_tokens
848            .checked_add(1)
849            .context("workload emitted output token count overflow")?;
850        Ok(())
851    }
852
853    pub fn on_complete(&mut self, request_uuid: Uuid, now_ms: f64) -> Result<()> {
854        self.on_terminal(request_uuid, now_ms, false)
855    }
856
857    pub fn on_terminal(&mut self, request_uuid: Uuid, now_ms: f64, rejected: bool) -> Result<()> {
858        let outcome = if rejected {
859            TurnOutcome::Rejected
860        } else {
861            TurnOutcome::Completed
862        };
863        let resolution = self
864            .resolve_turn(request_uuid, now_ms, outcome)?
865            .expect("completed turns require an in-flight request");
866        self.apply_resolution(resolution, now_ms);
867        Ok(())
868    }
869
870    fn resolve_turn(
871        &mut self,
872        request_uuid: Uuid,
873        now_ms: f64,
874        outcome: TurnOutcome,
875    ) -> Result<Option<TurnResolution>> {
876        let Some(in_flight) = self.in_flight.get(&request_uuid).copied() else {
877            return match outcome {
878                TurnOutcome::Completed | TurnOutcome::Rejected => Err(anyhow!(
879                    "unknown workload request completion for {request_uuid}"
880                )),
881                TurnOutcome::Cancelled => Ok(None),
882            };
883        };
884        let session = self
885            .sessions
886            .get(in_flight.session_index)
887            .ok_or_else(|| anyhow!("unknown workload session {}", in_flight.session_index))?;
888        let turn = session.turns.get(in_flight.turn_index).ok_or_else(|| {
889            anyhow!(
890                "unknown workload turn {} for session {}",
891                in_flight.turn_index,
892                session.session_id
893            )
894        })?;
895        if session.in_flight != Some(request_uuid) {
896            bail!(
897                "session {} resolution for {} does not match in-flight request {:?}",
898                session.session_id,
899                request_uuid,
900                session.in_flight
901            );
902        }
903        if session.next_turn_index != in_flight.turn_index {
904            bail!(
905                "session {} resolution for turn {} does not match next turn {}",
906                session.session_id,
907                in_flight.turn_index,
908                session.next_turn_index
909            );
910        }
911
912        let request_id = turn.request_id.clone();
913        if outcome == TurnOutcome::Rejected && in_flight.emitted_output_tokens != 0 {
914            bail!(
915                "rejected workload request {request_uuid} emitted {} output tokens",
916                in_flight.emitted_output_tokens
917            );
918        }
919        let completed_output_tokens = (outcome == TurnOutcome::Completed
920            && self.prompt_mode == PromptMode::DeltaCumulative)
921            .then(|| {
922                let planned_output_tokens = turn
923                    .output_token_ids
924                    .as_ref()
925                    .expect("delta turns must have planned output tokens");
926                planned_output_tokens[..in_flight.emitted_output_tokens].to_vec()
927            });
928        let (next_turn_index, next_ready_at_ms, session_ended) = match outcome {
929            TurnOutcome::Completed | TurnOutcome::Rejected => {
930                let next_turn_index = in_flight
931                    .turn_index
932                    .checked_add(1)
933                    .context("workload turn index overflow")?;
934                let has_more_turns = self.policy.schedules_sequential_turns()
935                    && next_turn_index < session.turns.len();
936                let next_ready_at_ms = has_more_turns
937                    .then(|| now_ms + session.turns[next_turn_index].delay_after_previous_ms);
938                (next_turn_index, next_ready_at_ms, !has_more_turns)
939            }
940            TurnOutcome::Cancelled => (session.turns.len(), None, true),
941        };
942
943        self.in_flight
944            .remove(&request_uuid)
945            .expect("validated in-flight request must still exist");
946        let session = &mut self.sessions[in_flight.session_index];
947        session.in_flight = None;
948        session.next_turn_index = next_turn_index;
949        session.next_ready_at_ms = next_ready_at_ms;
950        if next_ready_at_ms.is_some()
951            && let Some(output_tokens) = completed_output_tokens
952        {
953            session.cumulative_tokens.extend(output_tokens);
954        }
955        if let Some(ready_at_ms) = next_ready_at_ms {
956            self.ready_sessions.push(ReadySession {
957                ready_at_ms,
958                session_index: in_flight.session_index,
959                turn_index: next_turn_index,
960            });
961        }
962
963        Ok(Some(TurnResolution {
964            request_id,
965            session_ended,
966        }))
967    }
968
969    fn apply_resolution(&mut self, resolution: TurnResolution, now_ms: f64) {
970        match &mut self.policy {
971            SchedulingPolicy::Trace => {}
972            SchedulingPolicy::Concurrency(state) => {
973                if resolution.session_ended {
974                    state.on_session_finished(&mut self.sessions, &mut self.ready_sessions, now_ms);
975                }
976            }
977            SchedulingPolicy::Agentic(state) => {
978                if let Some(request_id) = resolution.request_id {
979                    state.release_dependents(
980                        &mut self.sessions,
981                        &mut self.ready_sessions,
982                        &request_id,
983                        now_ms,
984                    );
985                }
986            }
987        }
988    }
989
990    pub fn next_ready_time_ms(&mut self) -> Option<f64> {
991        if self.policy.at_dispatch_capacity(self.in_flight.len()) {
992            return None;
993        }
994        loop {
995            let ready_session = *self.ready_sessions.peek()?;
996            let session = &self.sessions[ready_session.session_index];
997            if session.in_flight.is_some()
998                || session.next_turn_index != ready_session.turn_index
999                || session.next_ready_at_ms != Some(ready_session.ready_at_ms)
1000            {
1001                self.ready_sessions.pop();
1002                continue;
1003            }
1004            return Some(ready_session.ready_at_ms);
1005        }
1006    }
1007
1008    pub fn is_drained(&self) -> bool {
1009        self.in_flight.is_empty()
1010            && self
1011                .sessions
1012                .iter()
1013                .all(|session| session.next_turn_index >= session.turns.len())
1014    }
1015
1016    pub fn total_turns(&self) -> usize {
1017        self.sessions
1018            .iter()
1019            .map(|session| session.turns.len())
1020            .sum()
1021    }
1022}
1023
1024#[cfg(test)]
1025mod tests {
1026    use super::*;
1027    use crate::loadgen::{AgenticTrace, AgenticTurnTrace, SessionTrace, Trace, TurnTrace};
1028
1029    fn assert_deterministic_output_plan(
1030        mut first_driver: WorkloadDriver,
1031        mut second_driver: WorkloadDriver,
1032        expected_len: usize,
1033    ) {
1034        let first = first_driver.pop_ready(0.0, usize::MAX);
1035        let second = second_driver.pop_ready(0.0, usize::MAX);
1036
1037        assert_eq!(first.len(), 1);
1038        assert_eq!(second.len(), 1);
1039        assert_eq!(
1040            first[0].request.output_token_ids,
1041            second[0].request.output_token_ids
1042        );
1043        assert_eq!(
1044            first[0].request.output_token_ids.as_ref().map(Vec::len),
1045            Some(expected_len)
1046        );
1047    }
1048
1049    #[test]
1050    fn hash_free_admission_preserves_request_without_router_metadata() {
1051        let trace = Trace {
1052            block_size: 2,
1053            sessions: vec![SessionTrace {
1054                session_id: "a".into(),
1055                first_arrival_timestamp_ms: Some(0.0),
1056                turns: vec![TurnTrace {
1057                    input_length: 4,
1058                    max_output_tokens: 1,
1059                    hash_ids: vec![10, 11],
1060                    ..Default::default()
1061                }],
1062            }],
1063        };
1064        let mut with_hashes = WorkloadDriver::new_trace(trace.clone(), 2).unwrap();
1065        let mut without_hashes =
1066            WorkloadDriver::new_trace_without_replay_hashes(trace, 2, false).unwrap();
1067
1068        let with_hashes = with_hashes.pop_ready(0.0, 1).pop().unwrap();
1069        let without_hashes = without_hashes.pop_ready(0.0, 1).pop().unwrap();
1070
1071        assert!(with_hashes.replay_hashes.is_some());
1072        assert!(without_hashes.replay_hashes.is_none());
1073        assert_eq!(without_hashes.request.tokens, with_hashes.request.tokens);
1074        assert_eq!(
1075            without_hashes.request.output_token_ids,
1076            with_hashes.request.output_token_ids
1077        );
1078    }
1079
1080    fn two_session_trace() -> Trace {
1081        Trace {
1082            block_size: 1,
1083            sessions: vec![
1084                SessionTrace {
1085                    session_id: "a".into(),
1086                    first_arrival_timestamp_ms: Some(0.0),
1087                    turns: vec![
1088                        TurnTrace {
1089                            input_length: 2,
1090                            max_output_tokens: 1,
1091                            hash_ids: vec![1, 2],
1092                            delay_after_previous_ms: 0.0,
1093                            ..Default::default()
1094                        },
1095                        TurnTrace {
1096                            input_length: 2,
1097                            max_output_tokens: 1,
1098                            hash_ids: vec![3, 4],
1099                            delay_after_previous_ms: 5.0,
1100                            ..Default::default()
1101                        },
1102                    ],
1103                },
1104                SessionTrace {
1105                    session_id: "b".into(),
1106                    first_arrival_timestamp_ms: Some(0.0),
1107                    turns: vec![TurnTrace {
1108                        input_length: 2,
1109                        max_output_tokens: 1,
1110                        hash_ids: vec![5, 6],
1111                        delay_after_previous_ms: 0.0,
1112                        ..Default::default()
1113                    }],
1114                },
1115            ],
1116        }
1117    }
1118
1119    /// A: 2 turns (turn-1 has a 5ms think-time). B, C: 1 turn each. Used for the cap>1
1120    /// transition / cancellation tests (w/ a third session pending behind a cap of 2).
1121    fn three_session_trace() -> Trace {
1122        let mut trace = two_session_trace();
1123        trace.sessions.push(SessionTrace {
1124            session_id: "c".into(),
1125            first_arrival_timestamp_ms: Some(0.0),
1126            turns: vec![TurnTrace {
1127                input_length: 2,
1128                max_output_tokens: 1,
1129                hash_ids: vec![7, 8],
1130                delay_after_previous_ms: 0.0,
1131                ..Default::default()
1132            }],
1133        });
1134        trace
1135    }
1136
1137    #[test]
1138    fn full_prompts_remain_deferred_until_dispatch() {
1139        let mut driver = WorkloadDriver::new_trace(two_session_trace(), 1).unwrap();
1140
1141        assert!(driver.sessions.iter().all(|session| {
1142            session
1143                .turns
1144                .iter()
1145                .all(|turn| matches!(turn.prompt_tokens, PromptTokens::Deferred { .. }))
1146        }));
1147
1148        let ready = driver.pop_ready(0.0, 1);
1149        assert_eq!(ready.len(), 1);
1150        assert_eq!(ready[0].request.tokens, vec![1, 2]);
1151        assert!(ready[0].replay_hashes.is_some());
1152    }
1153
1154    #[test]
1155    fn compact_dispatch_does_not_retain_materialized_prompt() {
1156        let mut driver = WorkloadDriver::new_trace(two_session_trace(), 1).unwrap();
1157
1158        let mut ready = driver.pop_ready_compact(0.0, 1);
1159
1160        assert_eq!(ready.len(), 1);
1161        let request = ready.pop().expect("one compact request").request;
1162        assert_eq!(request.input_length(), 2);
1163        assert!(request.metadata().tokens.is_empty());
1164        assert!(request.materialized_tokens().is_none());
1165        assert_eq!(request.into_direct_request().tokens, vec![1, 2]);
1166    }
1167
1168    #[test]
1169    fn delta_cumulative_prompts_remain_materialized_during_setup() {
1170        let driver =
1171            WorkloadDriver::new_concurrency_accumulating_deltas(two_session_trace(), 1, 1).unwrap();
1172
1173        assert!(driver.sessions.iter().all(|session| {
1174            session
1175                .turns
1176                .iter()
1177                .all(|turn| matches!(turn.prompt_tokens, PromptTokens::Materialized(_)))
1178        }));
1179    }
1180
1181    #[test]
1182    fn deferred_prompt_validation_preserves_setup_errors() {
1183        let trace = Trace {
1184            block_size: 4,
1185            sessions: vec![SessionTrace {
1186                session_id: "invalid".into(),
1187                first_arrival_timestamp_ms: Some(0.0),
1188                turns: vec![TurnTrace {
1189                    input_length: 5,
1190                    max_output_tokens: 1,
1191                    hash_ids: vec![1],
1192                    ..Default::default()
1193                }],
1194            }],
1195        };
1196
1197        let error = WorkloadDriver::new_trace(trace, 4).unwrap_err();
1198        assert!(
1199            error
1200                .to_string()
1201                .contains("input_length 5 exceeds synthesized capacity 4")
1202        );
1203    }
1204
1205    #[test]
1206    fn unknown_completion_preserves_in_flight_state() {
1207        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1208        let admitted = driver.pop_ready(0.0, usize::MAX);
1209        let request_uuid = admitted[0].request_uuid;
1210        let session_index = driver.in_flight[&request_uuid].session_index;
1211
1212        let error = driver.on_complete(Uuid::new_v4(), 1.0).unwrap_err();
1213
1214        assert!(
1215            error
1216                .to_string()
1217                .contains("unknown workload request completion")
1218        );
1219        assert!(driver.in_flight.contains_key(&request_uuid));
1220        assert_eq!(driver.sessions[session_index].in_flight, Some(request_uuid));
1221    }
1222
1223    #[test]
1224    fn unknown_cancellation_is_noop() {
1225        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1226        let admitted = driver.pop_ready(0.0, usize::MAX);
1227        let request_uuid = admitted[0].request_uuid;
1228        let session_index = driver.in_flight[&request_uuid].session_index;
1229
1230        driver.release_cap_slot(Uuid::new_v4(), 1.0);
1231
1232        assert!(driver.in_flight.contains_key(&request_uuid));
1233        assert_eq!(driver.sessions[session_index].in_flight, Some(request_uuid));
1234    }
1235
1236    #[test]
1237    fn inconsistent_session_mapping_preserves_in_flight_entry() {
1238        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1239        let admitted = driver.pop_ready(0.0, usize::MAX);
1240        let request_uuid = admitted[0].request_uuid;
1241        let session_index = driver.in_flight[&request_uuid].session_index;
1242        driver.sessions[session_index].in_flight = Some(Uuid::new_v4());
1243
1244        let error = driver.on_complete(request_uuid, 1.0).unwrap_err();
1245
1246        assert!(
1247            error
1248                .to_string()
1249                .contains("does not match in-flight request")
1250        );
1251        assert!(driver.in_flight.contains_key(&request_uuid));
1252        assert_eq!(driver.sessions[session_index].next_turn_index, 0);
1253    }
1254
1255    #[test]
1256    fn cap_clamps_pop_ready_when_limit_is_unbounded() {
1257        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1258
1259        let first = driver.pop_ready(0.0, usize::MAX);
1260        assert_eq!(first.len(), 1);
1261        let second = driver.pop_ready(0.0, usize::MAX);
1262        assert!(
1263            second.is_empty(),
1264            "cap should block dispatch while slot is held"
1265        );
1266    }
1267
1268    #[test]
1269    fn pop_ready_admits_next_turn_after_on_complete() {
1270        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1271
1272        let admitted = driver.pop_ready(0.0, usize::MAX);
1273        assert_eq!(admitted.len(), 1);
1274        let uuid = admitted[0].request_uuid;
1275        driver.on_complete(uuid, 10.0).unwrap();
1276
1277        // next admitted turn is *this* session's turn-1
1278        // (ready at completion 10 + think-time 5 = 15)
1279        let next = driver.pop_ready(15.0, usize::MAX);
1280        assert_eq!(next.len(), 1);
1281        assert_eq!(next[0].turn_index, 1);
1282        assert_ne!(next[0].request_uuid, uuid);
1283    }
1284
1285    #[test]
1286    fn concurrency_is_depth_first_holding_slot_across_think_time() {
1287        // Session A: 2 turns (turn-1 has a 5ms think-time). Session B: 1 turn. cap = 1.
1288        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1289
1290        // A.turn0 admitted; B is pending (not activated — cap is 1).
1291        let a0 = driver.pop_ready(0.0, usize::MAX);
1292        assert_eq!(a0.len(), 1);
1293        assert_eq!(a0[0].turn_index, 0);
1294        let a0_uuid = a0[0].request_uuid;
1295        driver.on_complete(a0_uuid, 10.0).unwrap();
1296
1297        // During A's think-time (turn-1 ready at 10+5=15), B must NOT slip in: A holds the slot.
1298        assert!(
1299            driver.pop_ready(10.0, usize::MAX).is_empty(),
1300            "B must not be admitted while A holds its slot in think-time"
1301        );
1302
1303        // A.turn1 dispatches before B ever starts (depth-first).
1304        let a1 = driver.pop_ready(15.0, usize::MAX);
1305        assert_eq!(a1.len(), 1);
1306        assert_eq!(a1[0].turn_index, 1);
1307        driver.on_complete(a1[0].request_uuid, 20.0).unwrap();
1308
1309        // Only now that A is fully done is B activated.
1310        let b0 = driver.pop_ready(20.0, usize::MAX);
1311        assert_eq!(b0.len(), 1);
1312        assert_eq!(b0[0].turn_index, 0);
1313        assert_ne!(b0[0].request_uuid, a0_uuid);
1314        assert!(!driver.is_drained(), "B still in flight");
1315        driver.on_complete(b0[0].request_uuid, 30.0).unwrap();
1316        assert!(driver.is_drained());
1317    }
1318
1319    #[test]
1320    fn concurrency_cap2_admits_pending_when_active_session_finishes() {
1321        // cap = 2: A (2 turns) and B (1 turn) start active; C (1 turn) is pending.
1322        let mut driver = WorkloadDriver::new_concurrency(three_session_trace(), 1, 2).unwrap();
1323
1324        // Initial cohort: A.t0 and B.t0 (the cap-2 set); C stays pending.
1325        let first = driver.pop_ready(0.0, usize::MAX);
1326        let mut ids: Vec<&str> = first.iter().map(|r| r.session_id.as_str()).collect();
1327        ids.sort();
1328        assert_eq!(
1329            ids,
1330            vec!["a", "b"],
1331            "cap-2 admits exactly A and B; C pending"
1332        );
1333        let a0 = first
1334            .iter()
1335            .find(|r| r.session_id == "a")
1336            .unwrap()
1337            .request_uuid;
1338        let b0 = first
1339            .iter()
1340            .find(|r| r.session_id == "b")
1341            .unwrap()
1342            .request_uuid;
1343
1344        // A finishes turn-0 → enters think-time (A.t1 ready at 10+5=15); A keeps its slot.
1345        driver.on_complete(a0, 10.0).unwrap();
1346        // B finishes its only turn → frees a slot → C is activated.
1347        driver.on_complete(b0, 10.0).unwrap();
1348
1349        // At t=10 only C is admittable (its freed slot); A is mid-think-time and retains
1350        // its slot — neither dropped nor re-admitted early.
1351        let at_10 = driver.pop_ready(10.0, usize::MAX);
1352        assert_eq!(at_10.len(), 1, "only C is admittable at t=10");
1353        assert_eq!(at_10[0].session_id, "c");
1354        assert_eq!(at_10[0].turn_index, 0);
1355
1356        // A's retained slot resumes once its think-time elapses (t=15), proving it was
1357        // never evicted by C's admission.
1358        let at_15 = driver.pop_ready(15.0, usize::MAX);
1359        assert_eq!(at_15.len(), 1);
1360        assert_eq!(
1361            (at_15[0].session_id.as_str(), at_15[0].turn_index),
1362            ("a", 1)
1363        );
1364    }
1365
1366    #[test]
1367    fn release_cap_slot_terminates_inflight_session_and_admits_pending() {
1368        // Mirrors an online InFlightGuard drop (cancellation), which calls release_cap_slot.
1369        // cap = 2: A (2 turns) + B (1 turn) active, C (1 turn) pending. A is in think-time,
1370        // B is in flight and gets cancelled.
1371        let mut driver = WorkloadDriver::new_concurrency(three_session_trace(), 1, 2).unwrap();
1372
1373        let first = driver.pop_ready(0.0, usize::MAX);
1374        let a0 = first
1375            .iter()
1376            .find(|r| r.session_id == "a")
1377            .unwrap()
1378            .request_uuid;
1379        let b0 = first
1380            .iter()
1381            .find(|r| r.session_id == "b")
1382            .unwrap()
1383            .request_uuid;
1384
1385        // A → think-time (A.t1 ready at 15), retains its slot.
1386        driver.on_complete(a0, 10.0).unwrap();
1387        // B cancelled in flight: the online guard drop releases B's slot and terminates it.
1388        driver.release_cap_slot(b0, 10.0);
1389
1390        // B's freed slot admits C; A's continuation is untouched.
1391        let at_10 = driver.pop_ready(10.0, usize::MAX);
1392        assert_eq!(at_10.len(), 1);
1393        assert_eq!(
1394            at_10[0].session_id, "c",
1395            "C admitted into the slot freed by B's cancellation"
1396        );
1397        driver.on_complete(at_10[0].request_uuid, 12.0).unwrap();
1398
1399        // A's continuation survived the cancellation and resumes after its think-time.
1400        let a1 = driver.pop_ready(15.0, usize::MAX);
1401        assert_eq!(a1.len(), 1);
1402        assert_eq!((a1[0].session_id.as_str(), a1[0].turn_index), ("a", 1));
1403        driver.on_complete(a1[0].request_uuid, 20.0).unwrap();
1404
1405        // A (2 turns), B (cancelled/terminated), C (1 turn) all resolved → drained.
1406        assert!(driver.is_drained());
1407    }
1408
1409    #[test]
1410    fn next_ready_time_ms_returns_none_at_cap() {
1411        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1412
1413        let admitted = driver.pop_ready(0.0, usize::MAX);
1414        assert_eq!(admitted.len(), 1);
1415
1416        assert!(
1417            driver.next_ready_time_ms().is_none(),
1418            "expected None while at cap even with ready sessions queued"
1419        );
1420
1421        driver.on_complete(admitted[0].request_uuid, 10.0).unwrap();
1422        assert!(
1423            driver.next_ready_time_ms().is_some(),
1424            "expected readiness after a slot is freed"
1425        );
1426    }
1427
1428    #[test]
1429    fn uncapped_concurrency_admits_all_sessions_up_to_caller_limit() {
1430        // usize::MAX cap == effectively uncapped: every session is activated, so the
1431        // caller's pop_ready limit is the only bound.
1432        let mut driver =
1433            WorkloadDriver::new_concurrency(two_session_trace(), 1, usize::MAX).unwrap();
1434
1435        let admitted = driver.pop_ready(0.0, 5);
1436        assert_eq!(
1437            admitted.len(),
1438            2,
1439            "both sessions should admit when uncapped"
1440        );
1441        assert!(driver.next_ready_time_ms().is_none());
1442    }
1443
1444    #[test]
1445    fn release_cap_slot_is_noop_after_on_complete() {
1446        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1447
1448        let admitted = driver.pop_ready(0.0, usize::MAX);
1449        let uuid = admitted[0].request_uuid;
1450        driver.on_complete(uuid, 5.0).unwrap();
1451
1452        // release_cap_slot after on_complete is a no-op (the in-flight entry is already
1453        // gone), so it must NOT double-decrement active_sessions. The session still holds its
1454        // slot for turn-1 (ready at 5 + think-time 5 = 10)
1455        driver.release_cap_slot(uuid, 5.0);
1456
1457        let next = driver.pop_ready(10.0, usize::MAX);
1458        assert_eq!(next.len(), 1);
1459        assert_eq!(next[0].turn_index, 1);
1460        assert_ne!(next[0].request_uuid, uuid);
1461    }
1462
1463    #[test]
1464    fn release_cap_slot_recovers_cap_when_on_complete_was_skipped() {
1465        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1466
1467        let admitted = driver.pop_ready(0.0, usize::MAX);
1468        assert_eq!(admitted.len(), 1);
1469
1470        driver.release_cap_slot(admitted[0].request_uuid, 0.0);
1471
1472        let next = driver.pop_ready(0.0, usize::MAX);
1473        assert_eq!(
1474            next.len(),
1475            1,
1476            "cap slot should be available after release_cap_slot"
1477        );
1478    }
1479
1480    #[test]
1481    fn release_cap_slot_terminates_session_so_is_drained_completes() {
1482        let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1483
1484        let admitted = driver.pop_ready(0.0, usize::MAX);
1485        assert_eq!(admitted.len(), 1);
1486        let stuck_uuid = admitted[0].request_uuid;
1487
1488        driver.release_cap_slot(stuck_uuid, 0.0);
1489
1490        let neighbor = driver.pop_ready(0.0, usize::MAX);
1491        assert_eq!(
1492            neighbor.len(),
1493            1,
1494            "other session must still be admissible after its neighbor was terminated"
1495        );
1496        driver.on_complete(neighbor[0].request_uuid, 1.0).unwrap();
1497
1498        assert!(
1499            driver.is_drained(),
1500            "is_drained must become true so run_workload can exit"
1501        );
1502    }
1503
1504    #[test]
1505    fn full_prompt_modes_plan_missing_output_token_ids_deterministically() {
1506        let trace = Trace {
1507            block_size: 1,
1508            sessions: vec![SessionTrace {
1509                session_id: "a".into(),
1510                first_arrival_timestamp_ms: Some(0.0),
1511                turns: vec![TurnTrace {
1512                    input_length: 2,
1513                    max_output_tokens: 3,
1514                    hash_ids: vec![10, 11],
1515                    ..Default::default()
1516                }],
1517            }],
1518        };
1519        assert_deterministic_output_plan(
1520            WorkloadDriver::new_trace(trace.clone(), 1).unwrap(),
1521            WorkloadDriver::new_trace(trace, 1).unwrap(),
1522            3,
1523        );
1524
1525        let trace = AgenticTrace {
1526            block_size: 1,
1527            turns: vec![AgenticTurnTrace {
1528                request_id: "r1".into(),
1529                session_id: "a".into(),
1530                input_length: 2,
1531                max_output_tokens: 3,
1532                hash_ids: vec![10, 11],
1533                first_ready_timestamp_ms: Some(0.0),
1534                prefix_reset: true,
1535                ..Default::default()
1536            }],
1537        };
1538        assert_deterministic_output_plan(
1539            WorkloadDriver::new_agentic_trace(trace.clone(), 1).unwrap(),
1540            WorkloadDriver::new_agentic_trace(trace, 1).unwrap(),
1541            3,
1542        );
1543    }
1544
1545    #[test]
1546    fn accumulating_delta_mode_includes_previous_output_tokens() {
1547        let trace = Trace {
1548            block_size: 4,
1549            sessions: vec![SessionTrace {
1550                session_id: "a".into(),
1551                first_arrival_timestamp_ms: Some(0.0),
1552                turns: vec![
1553                    TurnTrace {
1554                        input_length: 6,
1555                        max_output_tokens: 2,
1556                        output_token_ids: Some(vec![20, 21]),
1557                        replay_key: None,
1558                        hash_ids: vec![10, 11],
1559                        delay_after_previous_ms: 0.0,
1560                        priority: 3,
1561                        strict_priority: 4,
1562                        policy_class: None,
1563                    },
1564                    TurnTrace {
1565                        input_length: 3,
1566                        max_output_tokens: 1,
1567                        output_token_ids: None,
1568                        replay_key: None,
1569                        hash_ids: vec![12],
1570                        delay_after_previous_ms: 5.0,
1571                        priority: -2,
1572                        strict_priority: 7,
1573                        policy_class: None,
1574                    },
1575                ],
1576            }],
1577        };
1578        let mut driver = WorkloadDriver::new_concurrency_accumulating_deltas(trace, 4, 1).unwrap();
1579
1580        let first = driver.pop_ready(0.0, usize::MAX);
1581        assert_eq!(first.len(), 1);
1582        assert_eq!(first[0].request.tokens, vec![10, 10, 10, 10, 11, 11]);
1583        assert_eq!(first[0].request.output_token_ids, Some(vec![20, 21]));
1584        assert_eq!(first[0].request.priority, 3);
1585        assert_eq!(first[0].request.strict_priority, 4);
1586        driver.on_output_token(first[0].request_uuid, 20).unwrap();
1587        driver.on_output_token(first[0].request_uuid, 21).unwrap();
1588        driver.on_complete(first[0].request_uuid, 10.0).unwrap();
1589
1590        let second = driver.pop_ready(15.0, usize::MAX);
1591        assert_eq!(second.len(), 1);
1592        assert_eq!(
1593            second[0].request.tokens,
1594            vec![10, 10, 10, 10, 11, 11, 20, 21, 12, 12, 12]
1595        );
1596        assert_eq!(second[0].request.priority, -2);
1597        assert_eq!(second[0].request.strict_priority, 7);
1598    }
1599
1600    #[test]
1601    fn accumulating_delta_mode_plans_missing_output_token_ids() {
1602        let trace = Trace {
1603            block_size: 1,
1604            sessions: vec![SessionTrace {
1605                session_id: "a".into(),
1606                first_arrival_timestamp_ms: Some(0.0),
1607                turns: vec![
1608                    TurnTrace {
1609                        input_length: 2,
1610                        max_output_tokens: 3,
1611                        hash_ids: vec![10, 11],
1612                        ..Default::default()
1613                    },
1614                    TurnTrace {
1615                        input_length: 1,
1616                        max_output_tokens: 1,
1617                        hash_ids: vec![12],
1618                        ..Default::default()
1619                    },
1620                ],
1621            }],
1622        };
1623        let mut driver = WorkloadDriver::new_concurrency_accumulating_deltas(trace, 1, 1).unwrap();
1624
1625        let first = driver.pop_ready(0.0, usize::MAX);
1626        assert_eq!(first.len(), 1);
1627        let planned_output = first[0]
1628            .request
1629            .output_token_ids
1630            .clone()
1631            .expect("delta replay should plan synthetic outputs");
1632        assert_eq!(planned_output.len(), 3);
1633        for &token_id in &planned_output {
1634            driver
1635                .on_output_token(first[0].request_uuid, token_id)
1636                .unwrap();
1637        }
1638        driver.on_complete(first[0].request_uuid, 1.0).unwrap();
1639
1640        let second = driver.pop_ready(1.0, usize::MAX);
1641        assert_eq!(second.len(), 1);
1642        let mut expected = vec![10, 11];
1643        expected.extend(planned_output);
1644        expected.push(12);
1645        assert_eq!(second[0].request.tokens, expected);
1646        assert_eq!(
1647            second[0].request.output_token_ids.as_ref().map(Vec::len),
1648            Some(1)
1649        );
1650    }
1651
1652    #[test]
1653    fn accumulating_delta_mode_appends_only_emitted_output_tokens() {
1654        let trace = Trace {
1655            block_size: 1,
1656            sessions: vec![SessionTrace {
1657                session_id: "a".into(),
1658                first_arrival_timestamp_ms: Some(0.0),
1659                turns: vec![
1660                    TurnTrace {
1661                        input_length: 1,
1662                        max_output_tokens: 3,
1663                        output_token_ids: Some(vec![20, 21, 22]),
1664                        hash_ids: vec![10],
1665                        ..Default::default()
1666                    },
1667                    TurnTrace {
1668                        input_length: 1,
1669                        max_output_tokens: 1,
1670                        hash_ids: vec![12],
1671                        ..Default::default()
1672                    },
1673                ],
1674            }],
1675        };
1676        let mut driver = WorkloadDriver::new_concurrency_accumulating_deltas(trace, 1, 1).unwrap();
1677
1678        let first = driver.pop_ready(0.0, usize::MAX);
1679        driver.on_output_token(first[0].request_uuid, 20).unwrap();
1680        driver.on_output_token(first[0].request_uuid, 21).unwrap();
1681        driver.on_complete(first[0].request_uuid, 1.0).unwrap();
1682
1683        let second = driver.pop_ready(1.0, usize::MAX);
1684        assert_eq!(second[0].request.tokens, vec![10, 20, 21, 12]);
1685    }
1686
1687    #[test]
1688    fn accumulating_delta_mode_does_not_append_rejected_output_tokens() {
1689        let trace = Trace {
1690            block_size: 1,
1691            sessions: vec![SessionTrace {
1692                session_id: "a".into(),
1693                first_arrival_timestamp_ms: Some(0.0),
1694                turns: vec![
1695                    TurnTrace {
1696                        input_length: 1,
1697                        max_output_tokens: 2,
1698                        output_token_ids: Some(vec![20, 21]),
1699                        hash_ids: vec![10],
1700                        ..Default::default()
1701                    },
1702                    TurnTrace {
1703                        input_length: 1,
1704                        max_output_tokens: 1,
1705                        hash_ids: vec![12],
1706                        ..Default::default()
1707                    },
1708                ],
1709            }],
1710        };
1711        let mut driver = WorkloadDriver::new_concurrency_accumulating_deltas(trace, 1, 1).unwrap();
1712
1713        let first = driver.pop_ready(0.0, usize::MAX);
1714        driver
1715            .on_terminal(first[0].request_uuid, 1.0, true)
1716            .unwrap();
1717
1718        let second = driver.pop_ready(1.0, usize::MAX);
1719        assert_eq!(second[0].request.tokens, vec![10, 12]);
1720    }
1721
1722    #[test]
1723    fn agentic_mode_releases_turn_after_dependency_completion_plus_delay() {
1724        let trace = AgenticTrace {
1725            block_size: 1,
1726            turns: vec![
1727                AgenticTurnTrace {
1728                    request_id: "r1".into(),
1729                    session_id: "root".into(),
1730                    input_length: 2,
1731                    max_output_tokens: 1,
1732                    hash_ids: vec![1, 2],
1733                    first_ready_timestamp_ms: Some(0.0),
1734                    delay_after_dependencies_ms: 0.0,
1735                    wait_for: Vec::new(),
1736                    prefix_reset: true,
1737                    ..Default::default()
1738                },
1739                AgenticTurnTrace {
1740                    request_id: "r2".into(),
1741                    session_id: "root".into(),
1742                    input_length: 2,
1743                    max_output_tokens: 1,
1744                    hash_ids: vec![1, 3],
1745                    first_ready_timestamp_ms: Some(100.0),
1746                    delay_after_dependencies_ms: 5.0,
1747                    wait_for: vec!["r1".into()],
1748                    prefix_reset: false,
1749                    ..Default::default()
1750                },
1751            ],
1752        };
1753        let mut driver = WorkloadDriver::new_agentic_trace(trace, 1).unwrap();
1754
1755        let first = driver.pop_ready(0.0, usize::MAX);
1756        assert_eq!(first.len(), 1);
1757        assert_eq!(first[0].scheduled_ready_at_ms, 0.0);
1758        assert!(driver.pop_ready(14.0, usize::MAX).is_empty());
1759
1760        driver.on_complete(first[0].request_uuid, 10.0).unwrap();
1761        assert_eq!(driver.next_ready_time_ms(), Some(15.0));
1762        assert!(driver.pop_ready(14.0, usize::MAX).is_empty());
1763        let second = driver.pop_ready(15.0, usize::MAX);
1764        assert_eq!(second.len(), 1);
1765        assert_eq!(second[0].scheduled_ready_at_ms, 15.0);
1766    }
1767
1768    #[test]
1769    fn agentic_mode_releases_dependents_when_cap_slot_is_released() {
1770        let trace = AgenticTrace {
1771            block_size: 1,
1772            turns: vec![
1773                AgenticTurnTrace {
1774                    request_id: "r1".into(),
1775                    session_id: "root".into(),
1776                    input_length: 2,
1777                    max_output_tokens: 1,
1778                    hash_ids: vec![1, 2],
1779                    first_ready_timestamp_ms: Some(0.0),
1780                    delay_after_dependencies_ms: 0.0,
1781                    wait_for: Vec::new(),
1782                    prefix_reset: true,
1783                    ..Default::default()
1784                },
1785                AgenticTurnTrace {
1786                    request_id: "r2".into(),
1787                    session_id: "child".into(),
1788                    input_length: 2,
1789                    max_output_tokens: 1,
1790                    hash_ids: vec![1, 3],
1791                    first_ready_timestamp_ms: Some(100.0),
1792                    delay_after_dependencies_ms: 5.0,
1793                    wait_for: vec!["r1".into()],
1794                    prefix_reset: true,
1795                    ..Default::default()
1796                },
1797            ],
1798        };
1799        let mut driver = WorkloadDriver::new_agentic_trace(trace, 1).unwrap();
1800
1801        let first = driver.pop_ready(0.0, usize::MAX);
1802        assert_eq!(first.len(), 1);
1803
1804        driver.release_cap_slot(first[0].request_uuid, 10.0);
1805
1806        assert_eq!(driver.next_ready_time_ms(), Some(15.0));
1807        let second = driver.pop_ready(15.0, usize::MAX);
1808        assert_eq!(second.len(), 1);
1809        assert_eq!(second[0].scheduled_ready_at_ms, 15.0);
1810    }
1811
1812    #[test]
1813    fn agentic_mode_waits_for_slowest_dependency() {
1814        let trace = AgenticTrace {
1815            block_size: 1,
1816            turns: vec![
1817                AgenticTurnTrace {
1818                    request_id: "a".into(),
1819                    session_id: "a".into(),
1820                    input_length: 1,
1821                    max_output_tokens: 1,
1822                    hash_ids: vec![1],
1823                    first_ready_timestamp_ms: Some(0.0),
1824                    delay_after_dependencies_ms: 0.0,
1825                    wait_for: Vec::new(),
1826                    prefix_reset: true,
1827                    ..Default::default()
1828                },
1829                AgenticTurnTrace {
1830                    request_id: "b".into(),
1831                    session_id: "b".into(),
1832                    input_length: 1,
1833                    max_output_tokens: 1,
1834                    hash_ids: vec![2],
1835                    first_ready_timestamp_ms: Some(0.0),
1836                    delay_after_dependencies_ms: 0.0,
1837                    wait_for: Vec::new(),
1838                    prefix_reset: true,
1839                    ..Default::default()
1840                },
1841                AgenticTurnTrace {
1842                    request_id: "join".into(),
1843                    session_id: "root".into(),
1844                    input_length: 1,
1845                    max_output_tokens: 1,
1846                    hash_ids: vec![3],
1847                    first_ready_timestamp_ms: Some(1.0),
1848                    delay_after_dependencies_ms: 2.0,
1849                    wait_for: vec!["a".into(), "b".into()],
1850                    prefix_reset: false,
1851                    ..Default::default()
1852                },
1853            ],
1854        };
1855        let mut driver = WorkloadDriver::new_agentic_trace(trace, 1).unwrap();
1856
1857        let initial = driver.pop_ready(0.0, usize::MAX);
1858        assert_eq!(initial.len(), 2);
1859        driver.on_complete(initial[0].request_uuid, 10.0).unwrap();
1860        assert!(driver.next_ready_time_ms().is_none());
1861        driver.on_complete(initial[1].request_uuid, 30.0).unwrap();
1862        assert_eq!(driver.next_ready_time_ms(), Some(32.0));
1863    }
1864}