1use 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 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 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 #[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 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 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 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 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 let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1289
1290 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 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 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 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 let mut driver = WorkloadDriver::new_concurrency(three_session_trace(), 1, 2).unwrap();
1323
1324 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 driver.on_complete(a0, 10.0).unwrap();
1346 driver.on_complete(b0, 10.0).unwrap();
1348
1349 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 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 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 driver.on_complete(a0, 10.0).unwrap();
1387 driver.release_cap_slot(b0, 10.0);
1389
1390 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 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 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 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 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}