Skip to main content

kcode_telegram_update_dispatch/
lib.rs

1#![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}