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::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 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 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 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 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 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 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 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 let mut driver = WorkloadDriver::new_concurrency(two_session_trace(), 1, 1).unwrap();
1275
1276 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 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 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 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 let mut driver = WorkloadDriver::new_concurrency(three_session_trace(), 1, 2).unwrap();
1309
1310 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 driver.on_complete(a0, 10.0).unwrap();
1332 driver.on_complete(b0, 10.0).unwrap();
1334
1335 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 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 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 driver.on_complete(a0, 10.0).unwrap();
1373 driver.release_cap_slot(b0, 10.0);
1375
1376 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 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 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 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 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}