Skip to main content

aisimulate_core/replay/core/
round_robin.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::collections::BTreeMap;
5use std::marker::PhantomData;
6
7use anyhow::{Result, anyhow};
8use rustc_hash::FxHashMap;
9use uuid::Uuid;
10
11use super::{
12    EngineEventBatch, Placement, PlacementDecision, PlacementEffects, PlacementPolicy,
13    RequestIdentity, WorkerTopology,
14};
15
16#[derive(Debug)]
17pub struct AggregatedRoundRobin {
18    next_worker: usize,
19    next_rank_by_worker: FxHashMap<usize, u32>,
20    dp_size: u32,
21}
22
23#[derive(Debug)]
24pub struct AggregatedRoundRobinPlacement<Events: EngineEventBatch> {
25    counter: AggregatedRoundRobin,
26    workers: BTreeMap<usize, Vec<usize>>,
27    events: PhantomData<Events>,
28}
29
30impl<Events: EngineEventBatch> AggregatedRoundRobinPlacement<Events> {
31    pub fn new(dp_size: u32, workers: Vec<WorkerTopology>) -> Self {
32        let mut counter = AggregatedRoundRobin::new(dp_size);
33        for worker in &workers {
34            counter.worker_ready(worker.worker_id);
35        }
36        Self {
37            counter,
38            workers: workers
39                .into_iter()
40                .map(|worker| (worker.worker_id, worker.scheduler_ids))
41                .collect(),
42            events: PhantomData,
43        }
44    }
45}
46
47impl<Request, Events> PlacementPolicy<Request> for AggregatedRoundRobinPlacement<Events>
48where
49    Request: RequestIdentity,
50    Events: EngineEventBatch,
51{
52    type Metadata = ();
53    type Observation = Events;
54
55    #[inline]
56    fn place(
57        &mut self,
58        request: &Request,
59        _metadata: Self::Metadata,
60        _session_id: Option<String>,
61        _now_ms: f64,
62    ) -> Result<PlacementEffects> {
63        let request_id = request
64            .request_id()
65            .ok_or_else(|| anyhow!("round-robin placement requires a request UUID"))?;
66        let scheduler_id = self.counter.next(
67            self.workers.keys().copied(),
68            request.preferred_dp_rank(),
69            |worker_id, rank| {
70                self.workers
71                    .get(&worker_id)
72                    .and_then(|ranks| ranks.get(rank as usize))
73                    .copied()
74            },
75        )?;
76        Ok(PlacementEffects {
77            decision: PlacementDecision::Immediate(Placement {
78                request_id,
79                scheduler_id,
80                reported_overlap_tokens: 0,
81                cache_sample: None,
82                placement_replica_id: None,
83            }),
84            released: Vec::new(),
85        })
86    }
87
88    #[inline]
89    fn observe(&mut self, _observation: Events, _now_ms: f64) -> Result<Vec<Placement>> {
90        Ok(Vec::new())
91    }
92
93    #[inline]
94    fn cancel_pending(&mut self, _request_id: Uuid) -> bool {
95        false
96    }
97
98    #[inline]
99    fn request_terminal(&mut self, _request_id: Uuid, _now_ms: f64) -> Result<Vec<Placement>> {
100        Ok(Vec::new())
101    }
102
103    fn prefill_completed(&mut self, _request_id: Uuid, _now_ms: f64) -> Result<Vec<Placement>> {
104        Ok(Vec::new())
105    }
106
107    #[inline]
108    fn pending_count(&self) -> usize {
109        0
110    }
111
112    fn worker_ready(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
113        self.counter.worker_ready(worker.worker_id);
114        self.workers.insert(worker.worker_id, worker.scheduler_ids);
115        Ok(Vec::new())
116    }
117
118    fn worker_draining(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
119        self.workers.remove(&worker.worker_id);
120        Ok(Vec::new())
121    }
122
123    fn worker_removed(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
124        self.workers.remove(&worker.worker_id);
125        self.counter.worker_removed(worker.worker_id);
126        Ok(Vec::new())
127    }
128
129    #[inline]
130    fn topology_settled(&mut self, _now_ms: f64) -> Result<Vec<Placement>> {
131        Ok(Vec::new())
132    }
133}
134
135#[derive(Debug)]
136pub struct PoolRoundRobinPlacement<Events: EngineEventBatch> {
137    next: usize,
138    workers: BTreeMap<usize, Vec<usize>>,
139    events: PhantomData<Events>,
140}
141
142impl<Events: EngineEventBatch> PoolRoundRobinPlacement<Events> {
143    pub fn new(workers: Vec<WorkerTopology>) -> Self {
144        Self {
145            next: 0,
146            workers: workers
147                .into_iter()
148                .map(|worker| (worker.worker_id, worker.scheduler_ids))
149                .collect(),
150            events: PhantomData,
151        }
152    }
153}
154
155impl<Request, Events> PlacementPolicy<Request> for PoolRoundRobinPlacement<Events>
156where
157    Request: RequestIdentity,
158    Events: EngineEventBatch,
159{
160    type Metadata = ();
161    type Observation = Events;
162
163    fn place(
164        &mut self,
165        request: &Request,
166        _metadata: Self::Metadata,
167        _session_id: Option<String>,
168        _now_ms: f64,
169    ) -> Result<PlacementEffects> {
170        let request_id = request
171            .request_id()
172            .ok_or_else(|| anyhow!("round-robin placement requires a request UUID"))?;
173        let active_count = self.workers.values().map(Vec::len).sum::<usize>();
174        if active_count == 0 {
175            return Err(anyhow!("no active workers for round-robin placement"));
176        }
177        let index = self.next % active_count;
178        let scheduler_id = self
179            .workers
180            .values()
181            .flat_map(|ranks| ranks.iter().copied())
182            .nth(index)
183            .expect("active round-robin pool must contain a scheduler");
184        self.next = index + 1;
185        Ok(PlacementEffects {
186            decision: PlacementDecision::Immediate(Placement {
187                request_id,
188                scheduler_id,
189                reported_overlap_tokens: 0,
190                cache_sample: None,
191                placement_replica_id: None,
192            }),
193            released: Vec::new(),
194        })
195    }
196
197    fn observe(&mut self, _observation: Events, _now_ms: f64) -> Result<Vec<Placement>> {
198        Ok(Vec::new())
199    }
200
201    fn cancel_pending(&mut self, _request_id: Uuid) -> bool {
202        false
203    }
204
205    fn request_terminal(&mut self, _request_id: Uuid, _now_ms: f64) -> Result<Vec<Placement>> {
206        Ok(Vec::new())
207    }
208
209    fn prefill_completed(&mut self, _request_id: Uuid, _now_ms: f64) -> Result<Vec<Placement>> {
210        Ok(Vec::new())
211    }
212
213    fn pending_count(&self) -> usize {
214        0
215    }
216
217    fn worker_ready(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
218        self.workers.insert(worker.worker_id, worker.scheduler_ids);
219        Ok(Vec::new())
220    }
221
222    fn worker_draining(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
223        self.workers.remove(&worker.worker_id);
224        Ok(Vec::new())
225    }
226
227    fn worker_removed(&mut self, worker: WorkerTopology, _now_ms: f64) -> Result<Vec<Placement>> {
228        self.workers.remove(&worker.worker_id);
229        Ok(Vec::new())
230    }
231
232    fn topology_settled(&mut self, _now_ms: f64) -> Result<Vec<Placement>> {
233        Ok(Vec::new())
234    }
235}
236
237impl AggregatedRoundRobin {
238    pub fn new(dp_size: u32) -> Self {
239        Self {
240            next_worker: 0,
241            next_rank_by_worker: FxHashMap::default(),
242            dp_size: dp_size.max(1),
243        }
244    }
245
246    pub(crate) fn next(
247        &mut self,
248        mut active_workers: impl ExactSizeIterator<Item = usize>,
249        preferred_rank: Option<u32>,
250        rank_id: impl FnOnce(usize, u32) -> Option<usize>,
251    ) -> Result<usize> {
252        if active_workers.len() == 0 {
253            return Err(anyhow!("no active workers for round-robin placement"));
254        }
255        let index = self.next_worker % active_workers.len();
256        self.next_worker = index + 1;
257        let worker_id = active_workers
258            .nth(index)
259            .expect("active round-robin worker must exist at the selected index");
260        let rank = match preferred_rank {
261            Some(rank) if rank >= self.dp_size => {
262                return Err(anyhow!(
263                    "preferred attention-DP rank {rank} is out of range for dp_size {}",
264                    self.dp_size
265                ));
266            }
267            Some(rank) => rank,
268            None => {
269                let next_rank = self.next_rank_by_worker.entry(worker_id).or_default();
270                let rank = *next_rank % self.dp_size;
271                *next_rank = rank + 1;
272                rank
273            }
274        };
275        rank_id(worker_id, rank).ok_or_else(|| {
276            anyhow!("logical worker {worker_id} does not expose preferred attention-DP rank {rank}")
277        })
278    }
279
280    pub(crate) fn worker_removed(&mut self, worker_id: usize) {
281        self.next_rank_by_worker.remove(&worker_id);
282    }
283
284    fn worker_ready(&mut self, worker_id: usize) {
285        self.next_rank_by_worker.entry(worker_id).or_default();
286    }
287}
288
289#[cfg(test)]
290mod tests {
291    use super::*;
292
293    #[derive(Debug)]
294    struct TestRequest(Uuid);
295
296    impl RequestIdentity for TestRequest {
297        fn request_id(&self) -> Option<Uuid> {
298            Some(self.0)
299        }
300    }
301
302    #[derive(Debug)]
303    struct RankedTestRequest {
304        id: Uuid,
305        preferred_dp_rank: u32,
306    }
307
308    impl RequestIdentity for RankedTestRequest {
309        fn request_id(&self) -> Option<Uuid> {
310            Some(self.id)
311        }
312
313        fn preferred_dp_rank(&self) -> Option<u32> {
314            Some(self.preferred_dp_rank)
315        }
316    }
317
318    fn scheduler_id(policy: &mut PoolRoundRobinPlacement<()>, ordinal: u128) -> usize {
319        let effects = PlacementPolicy::<TestRequest>::place(
320            policy,
321            &TestRequest(Uuid::from_u128(ordinal)),
322            (),
323            None,
324            0.0,
325        )
326        .unwrap();
327        let PlacementDecision::Immediate(placement) = effects.decision else {
328            panic!("round-robin placement must be immediate");
329        };
330        placement.scheduler_id
331    }
332
333    #[test]
334    fn pool_rotation_preserves_position_after_topology_change() {
335        let mut policy = PoolRoundRobinPlacement::<()>::new(vec![
336            WorkerTopology {
337                worker_id: 0,
338                scheduler_ids: vec![10],
339            },
340            WorkerTopology {
341                worker_id: 1,
342                scheduler_ids: vec![11],
343            },
344            WorkerTopology {
345                worker_id: 2,
346                scheduler_ids: vec![12],
347            },
348        ]);
349
350        assert_eq!(
351            (1..=4)
352                .map(|ordinal| scheduler_id(&mut policy, ordinal))
353                .collect::<Vec<_>>(),
354            vec![10, 11, 12, 10]
355        );
356        PlacementPolicy::<TestRequest>::worker_draining(
357            &mut policy,
358            WorkerTopology {
359                worker_id: 2,
360                scheduler_ids: vec![12],
361            },
362            0.0,
363        )
364        .unwrap();
365
366        assert_eq!(scheduler_id(&mut policy, 5), 11);
367    }
368
369    #[test]
370    fn empty_pool_returns_an_error_instead_of_dividing_by_zero() {
371        let mut policy = PoolRoundRobinPlacement::<()>::new(Vec::new());
372        let error = PlacementPolicy::<TestRequest>::place(
373            &mut policy,
374            &TestRequest(Uuid::from_u128(1)),
375            (),
376            None,
377            0.0,
378        )
379        .unwrap_err();
380
381        assert!(error.to_string().contains("no active workers"));
382    }
383
384    #[test]
385    fn aggregated_round_robin_honors_authored_dp_rank_within_each_worker() {
386        let mut policy = AggregatedRoundRobinPlacement::<()>::new(
387            2,
388            vec![
389                WorkerTopology {
390                    worker_id: 0,
391                    scheduler_ids: vec![10, 11],
392                },
393                WorkerTopology {
394                    worker_id: 1,
395                    scheduler_ids: vec![20, 21],
396                },
397            ],
398        );
399
400        let place = |policy: &mut AggregatedRoundRobinPlacement<()>, ordinal, rank| {
401            let effects = PlacementPolicy::<RankedTestRequest>::place(
402                policy,
403                &RankedTestRequest {
404                    id: Uuid::from_u128(ordinal),
405                    preferred_dp_rank: rank,
406                },
407                (),
408                None,
409                0.0,
410            )?;
411            let PlacementDecision::Immediate(placement) = effects.decision else {
412                panic!("round-robin placement must be immediate");
413            };
414            Ok::<_, anyhow::Error>(placement.scheduler_id)
415        };
416
417        assert_eq!(place(&mut policy, 1, 1).unwrap(), 11);
418        assert_eq!(place(&mut policy, 2, 1).unwrap(), 21);
419        assert!(
420            place(&mut policy, 3, 2)
421                .unwrap_err()
422                .to_string()
423                .contains("out of range")
424        );
425    }
426}