Skip to main content

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