1use 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}