kcode_telegram_update_dispatch/
lib.rs1#![doc = include_str!("../Documentation.md")]
2#![forbid(unsafe_code)]
3
4use std::{
5 collections::{HashMap, VecDeque},
6 future::Future,
7 panic::AssertUnwindSafe,
8 sync::Arc,
9};
10
11use futures::{FutureExt, future::BoxFuture};
12use teloxide::types::{Message, Update, UpdateKind};
13use tokio::sync::{Mutex, Semaphore};
14
15const PER_PRINCIPAL_QUEUE_CAPACITY: usize = 32;
16const MAX_ACTIVE_PRINCIPAL_WORKERS: usize = 256;
17const GLOBAL_PROCESSING_CONCURRENCY: usize = 16;
18
19type UpdateProcessor = Arc<dyn Fn(Update) -> BoxFuture<'static, anyhow::Result<()>> + Send + Sync>;
20
21#[derive(Clone, Debug, Eq, Hash, PartialEq)]
22enum UpdateQueueKey {
23 Private(i64),
24 GroupUser(i64, i64),
25 GroupControl(i64),
26 Other,
27}
28
29struct DispatcherInner {
30 processor: UpdateProcessor,
31 workers: Mutex<HashMap<UpdateQueueKey, VecDeque<Update>>>,
32 concurrency: Arc<Semaphore>,
33}
34
35#[derive(Clone)]
36pub struct UpdateDispatcher {
37 inner: Arc<DispatcherInner>,
38}
39
40impl UpdateDispatcher {
41 pub fn new<F, Fut>(processor: F) -> Self
42 where
43 F: Fn(Update) -> Fut + Send + Sync + 'static,
44 Fut: Future<Output = anyhow::Result<()>> + Send + 'static,
45 {
46 let processor: UpdateProcessor = Arc::new(move |update| Box::pin(processor(update)));
47 Self {
48 inner: Arc::new(DispatcherInner {
49 processor,
50 workers: Mutex::new(HashMap::new()),
51 concurrency: Arc::new(Semaphore::new(GLOBAL_PROCESSING_CONCURRENCY)),
52 }),
53 }
54 }
55
56 pub async fn enqueue(&self, update: Update) -> anyhow::Result<()> {
57 let update_id = i64::from(update.id.0);
58 let key = update_queue_key(&update);
59 let spawn_worker = {
60 let mut workers = self.inner.workers.lock().await;
61 if let Some(queue) = workers.get_mut(&key) {
62 if queue.len() >= PER_PRINCIPAL_QUEUE_CAPACITY {
63 tracing::warn!(
64 update_id,
65 "dropping Telegram update because its principal queue is full"
66 );
67 return Ok(());
68 }
69 queue.push_back(update);
70 false
71 } else {
72 if workers.len() >= MAX_ACTIVE_PRINCIPAL_WORKERS {
73 tracing::warn!(
74 update_id,
75 active_principals = workers.len(),
76 "dropping Telegram update because the principal-worker limit is full"
77 );
78 return Ok(());
79 }
80 workers.insert(key.clone(), VecDeque::from([update]));
81 true
82 }
83 };
84
85 if spawn_worker {
86 tokio::spawn(run_worker(self.inner.clone(), key));
87 }
88 Ok(())
89 }
90}
91
92async fn remove_worker(inner: &DispatcherInner, key: &UpdateQueueKey) {
93 inner.workers.lock().await.remove(key);
94}
95
96async fn run_worker(inner: Arc<DispatcherInner>, key: UpdateQueueKey) {
97 loop {
98 let update = {
99 let mut workers = inner.workers.lock().await;
100 let Some(queue) = workers.get_mut(&key) else {
101 return;
102 };
103 if let Some(update) = queue.pop_front() {
104 update
105 } else {
106 workers.remove(&key);
107 return;
108 }
109 };
110
111 let update_id = i64::from(update.id.0);
112 let permit = match inner.concurrency.clone().acquire_owned().await {
113 Ok(permit) => permit,
114 Err(_) => {
115 remove_worker(&inner, &key).await;
116 return;
117 }
118 };
119 match AssertUnwindSafe((inner.processor)(update))
120 .catch_unwind()
121 .await
122 {
123 Ok(Ok(())) => {}
124 Ok(Err(_)) => {
125 tracing::warn!(
126 update_id,
127 "Telegram update failed locally; later work will continue"
128 );
129 }
130 Err(_) => {
131 tracing::error!(
132 update_id,
133 "Telegram update processor panicked; later work will continue"
134 );
135 }
136 }
137 drop(permit);
138 }
139}
140
141fn update_queue_key(update: &Update) -> UpdateQueueKey {
142 match &update.kind {
143 UpdateKind::Message(message) | UpdateKind::EditedMessage(message) => {
144 if message.chat.is_private() {
145 message
146 .from
147 .as_ref()
148 .and_then(|user| i64::try_from(user.id.0).ok())
149 .map(UpdateQueueKey::Private)
150 .unwrap_or(UpdateQueueKey::Other)
151 } else if message.chat.is_group() || message.chat.is_supergroup() {
152 if message.migrate_to_chat_id().is_some()
153 || message.migrate_from_chat_id().is_some()
154 || is_group_authored_message(message)
155 {
156 UpdateQueueKey::GroupControl(message.chat.id.0)
157 } else {
158 message
159 .from
160 .as_ref()
161 .and_then(|user| i64::try_from(user.id.0).ok())
162 .map(|user_id| UpdateQueueKey::GroupUser(message.chat.id.0, user_id))
163 .unwrap_or(UpdateQueueKey::GroupControl(message.chat.id.0))
164 }
165 } else {
166 UpdateQueueKey::Other
167 }
168 }
169 UpdateKind::ChatMember(change) | UpdateKind::MyChatMember(change) => {
170 UpdateQueueKey::GroupControl(change.chat.id.0)
171 }
172 _ => UpdateQueueKey::Other,
173 }
174}
175
176fn is_group_authored_message(message: &Message) -> bool {
177 message
178 .sender_chat
179 .as_ref()
180 .is_some_and(|sender| sender.id == message.chat.id)
181 || message
182 .from
183 .as_ref()
184 .is_some_and(|user| user.is_anonymous())
185}
186
187#[cfg(test)]
188mod tests {
189 use super::*;
190 use serde_json::json;
191 use std::sync::Mutex as StdMutex;
192 use teloxide::types::UpdateId;
193 use tokio::sync::{Semaphore, mpsc};
194
195 fn private_message(message_id: i32, user_id: i64, text: &str) -> Message {
196 serde_json::from_value(json!({
197 "message_id":message_id,
198 "date":1629404938,
199 "from":{"id":user_id,"is_bot":false,"first_name":"User"},
200 "chat":{"id":user_id,"first_name":"User","type":"private"},
201 "text":text
202 }))
203 .unwrap()
204 }
205
206 fn private_update(update_id: u32, user_id: i64) -> Update {
207 Update {
208 id: UpdateId(update_id),
209 kind: UpdateKind::Message(private_message(
210 i32::try_from(update_id).unwrap(),
211 user_id,
212 "test",
213 )),
214 }
215 }
216
217 fn edited_private_update(update_id: u32, user_id: i64) -> Update {
218 Update {
219 id: UpdateId(update_id),
220 kind: UpdateKind::EditedMessage(private_message(1, user_id, "edited")),
221 }
222 }
223
224 fn group_message(message_id: i32, chat_id: i64, user_id: i64, group_authored: bool) -> Message {
225 let mut value = json!({
226 "message_id":message_id,
227 "date":1629404938,
228 "chat":{"id":chat_id,"title":"Friends","type":"supergroup"},
229 "text":"test"
230 });
231 if group_authored {
232 value["sender_chat"] = json!({"id":chat_id,"title":"Friends","type":"supergroup"});
233 } else {
234 value["from"] = json!({"id":user_id,"is_bot":false,"first_name":"User"});
235 }
236 serde_json::from_value(value).unwrap()
237 }
238
239 fn group_update(update_id: u32, chat_id: i64, user_id: i64, group_authored: bool) -> Update {
240 Update {
241 id: UpdateId(update_id),
242 kind: UpdateKind::Message(group_message(
243 i32::try_from(update_id).unwrap(),
244 chat_id,
245 user_id,
246 group_authored,
247 )),
248 }
249 }
250
251 #[tokio::test]
252 async fn principals_overlap_while_each_principal_remains_ordered() {
253 let slow_gate = Arc::new(Semaphore::new(0));
254 let completed = Arc::new(StdMutex::new(Vec::new()));
255 let (started_sender, mut started_receiver) = mpsc::unbounded_channel();
256 let (completion_sender, mut completion_receiver) = mpsc::unbounded_channel();
257
258 let processor = {
259 let slow_gate = slow_gate.clone();
260 let completed = completed.clone();
261 move |update: Update| {
262 let slow_gate = slow_gate.clone();
263 let completed = completed.clone();
264 let started_sender = started_sender.clone();
265 let completion_sender = completion_sender.clone();
266 async move {
267 let update_id = i64::from(update.id.0);
268 if update_id == 1 {
269 let _ = started_sender.send(());
270 let permit = slow_gate.acquire().await?;
271 permit.forget();
272 }
273 completed.lock().unwrap().push(update_id);
274 let _ = completion_sender.send(update_id);
275 Ok(())
276 }
277 }
278 };
279 let dispatcher = UpdateDispatcher::new(processor);
280
281 dispatcher.enqueue(private_update(1, 42)).await.unwrap();
282 tokio::time::timeout(std::time::Duration::from_secs(1), started_receiver.recv())
283 .await
284 .unwrap()
285 .unwrap();
286
287 dispatcher.enqueue(private_update(3, 42)).await.unwrap();
288 dispatcher.enqueue(private_update(2, 77)).await.unwrap();
289
290 assert_eq!(
291 tokio::time::timeout(
292 std::time::Duration::from_secs(1),
293 completion_receiver.recv(),
294 )
295 .await
296 .unwrap(),
297 Some(2)
298 );
299
300 slow_gate.add_permits(1);
301 assert_eq!(
302 tokio::time::timeout(
303 std::time::Duration::from_secs(1),
304 completion_receiver.recv(),
305 )
306 .await
307 .unwrap(),
308 Some(1)
309 );
310 assert_eq!(
311 tokio::time::timeout(
312 std::time::Duration::from_secs(1),
313 completion_receiver.recv(),
314 )
315 .await
316 .unwrap(),
317 Some(3)
318 );
319 assert_eq!(*completed.lock().unwrap(), vec![2, 1, 3]);
320 }
321
322 #[tokio::test]
323 async fn processor_errors_are_local_to_one_update() {
324 let (completion_sender, mut completion_receiver) = mpsc::unbounded_channel();
325 let dispatcher = UpdateDispatcher::new(move |update: Update| {
326 let completion_sender = completion_sender.clone();
327 async move {
328 let update_id = i64::from(update.id.0);
329 if update_id == 1 {
330 anyhow::bail!("test processor error");
331 }
332 let _ = completion_sender.send(update_id);
333 Ok(())
334 }
335 });
336
337 dispatcher.enqueue(private_update(1, 42)).await.unwrap();
338 dispatcher.enqueue(private_update(2, 42)).await.unwrap();
339
340 assert_eq!(
341 tokio::time::timeout(
342 std::time::Duration::from_secs(1),
343 completion_receiver.recv(),
344 )
345 .await
346 .unwrap(),
347 Some(2)
348 );
349 }
350
351 #[tokio::test]
352 async fn processor_panics_are_local_to_one_update() {
353 let (completion_sender, mut completion_receiver) = mpsc::unbounded_channel();
354 let dispatcher = UpdateDispatcher::new(move |update: Update| {
355 let completion_sender = completion_sender.clone();
356 async move {
357 let update_id = i64::from(update.id.0);
358 if update_id == 1 {
359 panic!("test processor panic");
360 }
361 let _ = completion_sender.send(update_id);
362 Ok(())
363 }
364 });
365
366 dispatcher.enqueue(private_update(1, 42)).await.unwrap();
367 dispatcher.enqueue(private_update(2, 42)).await.unwrap();
368
369 assert_eq!(
370 tokio::time::timeout(
371 std::time::Duration::from_secs(1),
372 completion_receiver.recv(),
373 )
374 .await
375 .unwrap(),
376 Some(2)
377 );
378 }
379
380 #[tokio::test]
381 async fn a_drained_principal_is_removed_before_a_successor_worker_starts() {
382 let completed = Arc::new(StdMutex::new(Vec::new()));
383 let (completion_sender, mut completion_receiver) = mpsc::unbounded_channel();
384
385 let processor = {
386 let completed = completed.clone();
387 move |update: Update| {
388 let completed = completed.clone();
389 let completion_sender = completion_sender.clone();
390 async move {
391 let update_id = i64::from(update.id.0);
392 completed.lock().unwrap().push(update_id);
393 let _ = completion_sender.send(update_id);
394 Ok(())
395 }
396 }
397 };
398 let dispatcher = UpdateDispatcher::new(processor);
399
400 dispatcher.enqueue(private_update(1, 42)).await.unwrap();
401 assert_eq!(
402 tokio::time::timeout(
403 std::time::Duration::from_secs(1),
404 completion_receiver.recv(),
405 )
406 .await
407 .unwrap(),
408 Some(1)
409 );
410
411 loop {
412 if dispatcher.inner.workers.lock().await.is_empty() {
413 break;
414 }
415 tokio::task::yield_now().await;
416 }
417
418 dispatcher.enqueue(private_update(2, 42)).await.unwrap();
419 assert_eq!(
420 tokio::time::timeout(
421 std::time::Duration::from_secs(1),
422 completion_receiver.recv(),
423 )
424 .await
425 .unwrap(),
426 Some(2)
427 );
428 assert_eq!(*completed.lock().unwrap(), vec![1, 2]);
429 }
430
431 #[test]
432 fn edits_share_their_source_principals_queue() {
433 let original = private_update(10, 42);
434 let edited = edited_private_update(11, 42);
435
436 assert_eq!(update_queue_key(&original), update_queue_key(&edited));
437 assert_ne!(
438 update_queue_key(&original),
439 update_queue_key(&private_update(12, 77))
440 );
441 }
442
443 #[test]
444 fn group_user_and_control_updates_use_separate_keys() {
445 let ordinary = group_update(1, -100, 42, false);
446 let control = group_update(2, -100, 42, true);
447
448 assert_eq!(
449 update_queue_key(&ordinary),
450 UpdateQueueKey::GroupUser(-100, 42)
451 );
452 assert_eq!(
453 update_queue_key(&control),
454 UpdateQueueKey::GroupControl(-100)
455 );
456 }
457
458 #[test]
459 fn unkeyed_updates_share_one_bounded_worker_class() {
460 let first = Update {
461 id: UpdateId(1),
462 kind: UpdateKind::Error("unknown".into()),
463 };
464 let second = Update {
465 id: UpdateId(2),
466 kind: UpdateKind::Error("unknown".into()),
467 };
468
469 assert_eq!(update_queue_key(&first), UpdateQueueKey::Other);
470 assert_eq!(update_queue_key(&first), update_queue_key(&second));
471 }
472}