agora_agentkit/scheduler/
grouping.rs1use std::collections::BTreeMap;
7
8use super::WorkItem;
9
10#[derive(Debug, Clone)]
12pub struct GroupingConfig {
13 pub context_length_bucket: u32,
17}
18
19impl Default for GroupingConfig {
20 fn default() -> Self {
21 Self {
22 context_length_bucket: 4096,
23 }
24 }
25}
26
27#[derive(Debug)]
32pub struct BatchGroup<P> {
33 pub model: String,
35 pub prefix_hash: u64,
37 pub context_bucket: u32,
39 pub items: Vec<WorkItem<P>>,
41}
42
43type GroupKey = (String, u64, u32);
48
49pub fn group_work_items<P>(
55 items: Vec<WorkItem<P>>,
56 config: &GroupingConfig,
57) -> Vec<BatchGroup<P>> {
58 let bucket_size = config.context_length_bucket.max(1); let mut groups: BTreeMap<GroupKey, Vec<WorkItem<P>>> = BTreeMap::new();
61
62 for item in items {
63 let bucket = item.token_count / bucket_size;
64 let key = (item.model.clone(), item.prefix_hash, bucket);
65 groups.entry(key).or_default().push(item);
66 }
67
68 groups
69 .into_iter()
70 .map(|((model, prefix_hash, context_bucket), items)| BatchGroup {
71 model,
72 prefix_hash,
73 context_bucket,
74 items,
75 })
76 .collect()
77}
78
79#[cfg(test)]
80mod tests {
81 use std::time::Instant;
82
83 use crate::ids::AgentId;
84
85 use super::super::CycleStep;
86 use super::*;
87
88 fn item(model: &str, prefix_hash: u64, token_count: u32) -> WorkItem<()> {
89 WorkItem {
90 agent_id: AgentId::new(),
91 prompt: (),
92 step: CycleStep::Think,
93 prefix_hash,
94 model: model.to_string(),
95 queued_at: Instant::now(),
96 token_count,
97 }
98 }
99
100 #[test]
101 fn groups_by_model() {
102 let items = vec![
103 item("claude", 100, 5000),
104 item("cogito", 100, 5000),
105 item("claude", 100, 5000),
106 ];
107
108 let groups = group_work_items(items, &GroupingConfig::default());
109 assert_eq!(groups.len(), 2);
110 assert_eq!(groups[0].model, "claude");
111 assert_eq!(groups[0].items.len(), 2);
112 assert_eq!(groups[1].model, "cogito");
113 assert_eq!(groups[1].items.len(), 1);
114 }
115
116 #[test]
117 fn groups_by_prefix_hash() {
118 let items = vec![
119 item("claude", 100, 5000),
120 item("claude", 200, 5000),
121 item("claude", 100, 5000),
122 ];
123
124 let groups = group_work_items(items, &GroupingConfig::default());
125 assert_eq!(groups.len(), 2);
126 assert_eq!(groups[0].prefix_hash, 100);
127 assert_eq!(groups[0].items.len(), 2);
128 assert_eq!(groups[1].prefix_hash, 200);
129 }
130
131 #[test]
132 fn groups_by_context_bucket() {
133 let config = GroupingConfig {
134 context_length_bucket: 4096,
135 };
136
137 let items = vec![
138 item("claude", 100, 4000), item("claude", 100, 5000), item("claude", 100, 4500), item("claude", 100, 12000), ];
143
144 let groups = group_work_items(items, &config);
145 assert_eq!(groups.len(), 3);
146 assert_eq!(groups[0].context_bucket, 0);
147 assert_eq!(groups[0].items.len(), 1);
148 assert_eq!(groups[1].context_bucket, 1);
149 assert_eq!(groups[1].items.len(), 2);
150 assert_eq!(groups[2].context_bucket, 2);
151 assert_eq!(groups[2].items.len(), 1);
152 }
153
154 #[test]
155 fn empty_input() {
156 let groups = group_work_items::<()>(vec![], &GroupingConfig::default());
157 assert!(groups.is_empty());
158 }
159}