Skip to main content

bloop_server_framework/
engine.rs

1use crate::achievement::{Achievement, AchievementAwardBatch, AchievementContext, AwardedTracker};
2use crate::bloop::{Bloop, BloopProvider, ProcessedBloop, bloops_since};
3use crate::event::Event;
4use crate::message::{
5    AchievementRecord, AudioData, BloopAccepted, DataHash, ErrorResponse, PreloadMatch,
6    PreloadMismatch, ServerMessage,
7};
8use crate::nfc_uid::NfcUid;
9use crate::player::{PlayerInfo, PlayerMutator, PlayerRegistry};
10use crate::trigger::TriggerRegistry;
11use chrono::{DateTime, Utc};
12use serde::Deserialize;
13use std::collections::{HashMap, HashSet};
14use std::fmt::Debug;
15use std::path::{Path, PathBuf};
16use std::sync::Arc;
17use std::time::Duration;
18use thiserror::Error;
19use tokio::fs::File;
20use tokio::io::AsyncReadExt;
21use tokio::sync::{Mutex, broadcast, mpsc, oneshot};
22#[cfg(feature = "tokio-graceful-shutdown")]
23use tokio_graceful_shutdown::{FutureExt, IntoSubsystem, SubsystemHandle};
24use tracing::{info, instrument, warn};
25use uuid::Uuid;
26
27#[derive(Debug)]
28pub enum EngineRequest {
29    Bloop { client_id: String, nfc_uid: NfcUid },
30    RetrieveAudio { id: Uuid },
31    PreloadCheck { manifest_hash: Option<DataHash> },
32}
33
34#[derive(Debug, Deserialize)]
35pub struct Throttle {
36    max_bloops: usize,
37    #[serde(with = "humantime_serde")]
38    threshold: Duration,
39}
40
41impl Throttle {
42    pub fn new(max_bloops: usize, threshold: Duration) -> Self {
43        Self {
44            max_bloops,
45            threshold,
46        }
47    }
48}
49
50struct HotAchievement {
51    id: Uuid,
52    client_id: String,
53    until: DateTime<Utc>,
54}
55
56impl HotAchievement {
57    fn new<Player: PlayerInfo>(id: Uuid, bloop: &Bloop<Player>, duration: Duration) -> Self {
58        Self {
59            id,
60            client_id: bloop.client_id.clone(),
61            until: bloop.recorded_at + duration,
62        }
63    }
64}
65
66pub struct Engine<Metadata, Player, State, Trigger>
67where
68    Player: PlayerInfo + PlayerMutator,
69    Trigger: Copy,
70{
71    bloop_provider: BloopProvider<Player>,
72    achievements: HashMap<Uuid, Achievement<Metadata, Player, State, Trigger>>,
73    audio_base_path: PathBuf,
74    audio_manifest_hash: DataHash,
75    player_registry: Arc<Mutex<PlayerRegistry<Player>>>,
76    state: Arc<Mutex<State>>,
77    trigger_registry: TriggerRegistry<Trigger>,
78    hot_achievements: Vec<HotAchievement>,
79    network_rx: mpsc::Receiver<(EngineRequest, oneshot::Sender<ServerMessage>)>,
80    event_tx: broadcast::Sender<Event>,
81    throttle: Option<Throttle>,
82}
83
84impl<Metadata, Player, State, Trigger> Engine<Metadata, Player, State, Trigger>
85where
86    Player: PlayerInfo + PlayerMutator,
87    Trigger: Copy,
88{
89    pub async fn process_requests(&mut self) {
90        while let Some((request, response)) = self.network_rx.recv().await {
91            match request {
92                EngineRequest::Bloop { nfc_uid, client_id } => {
93                    self.handle_bloop(nfc_uid, client_id, response).await;
94                }
95                EngineRequest::RetrieveAudio { id } => {
96                    self.handle_retrieve_audio(id, response);
97                }
98                EngineRequest::PreloadCheck { manifest_hash } => {
99                    self.handle_preload_check(manifest_hash, response);
100                }
101            }
102        }
103    }
104
105    #[instrument(skip(self, response))]
106    async fn handle_bloop(
107        &mut self,
108        nfc_uid: NfcUid,
109        client_id: String,
110        response: oneshot::Sender<ServerMessage>,
111    ) {
112        if self
113            .trigger_registry
114            .try_activate_trigger(nfc_uid, &client_id)
115        {
116            let _ = response.send(
117                BloopAccepted {
118                    achievements: Vec::new(),
119                }
120                .into(),
121            );
122            return;
123        }
124
125        let player = {
126            let player_registry = self.player_registry.lock().await;
127            let Some(player) = player_registry.get_by_nfc_uid(nfc_uid) else {
128                let _ = response.send(ServerMessage::Error(ErrorResponse::UnknownNfcUid));
129                return;
130            };
131            player
132        };
133
134        if let Some(throttle) = self.throttle.as_ref() {
135            let player_id = player.read().unwrap().id();
136
137            let recent_bloops = self
138                .bloop_provider
139                .for_client(&client_id)
140                .iter()
141                .filter(bloops_since(Utc::now() - throttle.threshold))
142                .take(throttle.max_bloops)
143                .collect::<Vec<_>>();
144
145            if recent_bloops
146                .iter()
147                .all(|bloop| bloop.player_id == player_id)
148                && recent_bloops.len() == throttle.max_bloops
149            {
150                let _ = response.send(ServerMessage::Error(ErrorResponse::NfcUidThrottled));
151                return;
152            }
153        }
154
155        player.write().unwrap().increment_bloops();
156        let bloop = Bloop::new(player.clone(), client_id, Utc::now());
157
158        let mut awarded_tracker = self.evaluate_achievements(&bloop).await;
159        self.activate_hot_achievements(&bloop, &awarded_tracker);
160        self.inject_hot_achievements(&bloop, &mut awarded_tracker);
161
162        let player_registry = self.player_registry.lock().await;
163        awarded_tracker.remove_duplicates(player_registry);
164
165        let achievement_ids: Vec<Uuid> = awarded_tracker
166            .for_player(bloop.player_id)
167            .map_or_else(Vec::new, |set| set.iter().cloned().collect());
168
169        self.apply_awarded(awarded_tracker).await;
170
171        let processed_bloop: ProcessedBloop = (&bloop).into();
172        self.bloop_provider.add(Arc::new(bloop));
173        let _ = self.event_tx.send(Event::BloopProcessed(processed_bloop));
174
175        let _ = response.send(
176            BloopAccepted {
177                achievements: achievement_ids
178                    .into_iter()
179                    .map(|id| AchievementRecord {
180                        id,
181                        audio_hash: self
182                            .achievements
183                            .get(&id)
184                            .unwrap()
185                            .audio_file
186                            .resolve(&self.audio_base_path)
187                            .as_ref()
188                            .map(|file| file.hash.clone()),
189                    })
190                    .collect(),
191            }
192            .into(),
193        );
194    }
195
196    async fn evaluate_achievements(&mut self, bloop: &Bloop<Player>) -> AwardedTracker {
197        let previous_awarded: HashSet<Uuid> = {
198            let player = bloop.player();
199            player.awarded_achievements().keys().cloned().collect()
200        };
201        let state = self.state.lock().await;
202        let ctx = AchievementContext::new(
203            bloop,
204            &self.bloop_provider,
205            &*state,
206            &mut self.trigger_registry,
207        );
208
209        for achievement in self.achievements.values() {
210            if !previous_awarded.contains(&achievement.id) {
211                achievement.evaluate(&ctx);
212            }
213        }
214
215        ctx.take_awarded_tracker()
216    }
217
218    fn inject_hot_achievements(
219        &mut self,
220        bloop: &Bloop<Player>,
221        awarded_tracker: &mut AwardedTracker,
222    ) {
223        let awarded = awarded_tracker.for_player_mut(bloop.player_id);
224
225        self.hot_achievements.retain(|hot_achievement| {
226            if hot_achievement.until < bloop.recorded_at {
227                return false;
228            }
229
230            if hot_achievement.client_id == bloop.client_id {
231                awarded.insert(hot_achievement.id);
232            }
233
234            true
235        });
236    }
237
238    fn activate_hot_achievements(
239        &mut self,
240        bloop: &Bloop<Player>,
241        awarded_tracker: &AwardedTracker,
242    ) {
243        let Some(awarded) = awarded_tracker.for_player(bloop.player_id) else {
244            return;
245        };
246
247        for achievement_id in awarded {
248            let Some(achievement) = self.achievements.get(achievement_id) else {
249                continue;
250            };
251
252            if let Some(hot_duration) = achievement.hot_duration {
253                self.hot_achievements.push(HotAchievement::new(
254                    *achievement_id,
255                    bloop,
256                    hot_duration,
257                ));
258            };
259        }
260    }
261
262    async fn apply_awarded(&self, tracker: AwardedTracker) {
263        let mut player_registry = self.player_registry.lock().await;
264        let batch: AchievementAwardBatch = tracker.into();
265
266        for player_awards in batch.players.iter() {
267            player_registry.mutate_by_id(player_awards.player_id, |player| {
268                for achievement_id in player_awards.achievement_ids.iter() {
269                    player.add_awarded_achievement(*achievement_id, batch.awarded_at);
270                }
271            });
272        }
273
274        let _ = self.event_tx.send(Event::AchievementsAwarded(batch));
275    }
276
277    #[instrument(skip(self, response))]
278    fn handle_retrieve_audio(&self, id: Uuid, response: oneshot::Sender<ServerMessage>) {
279        let Some(achievement) = self.achievements.get(&id) else {
280            info!("Client requested unknown achievement: {}", id);
281            let _ = response.send(ServerMessage::Error(ErrorResponse::AudioUnavailable));
282            return;
283        };
284
285        let Some(audio_file) = achievement.audio_file.resolve(&self.audio_base_path) else {
286            info!("Client requested audio for audio-less achievement: {}", id);
287            let _ = response.send(ServerMessage::Error(ErrorResponse::AudioUnavailable));
288            return;
289        };
290
291        let path = audio_file.path.clone();
292
293        tokio::spawn(async move {
294            let mut file = match File::open(&path).await {
295                Ok(file) => file,
296                Err(err) => {
297                    warn!("Failed to open file {:?}: {:?}", &path, err);
298                    let _ = response.send(ServerMessage::Error(ErrorResponse::AudioUnavailable));
299                    return;
300                }
301            };
302
303            let mut data = vec![];
304
305            if let Err(err) = file.read_to_end(&mut data).await {
306                warn!("Failed to read file {:?}: {:?}", &path, err);
307                let _ = response.send(ServerMessage::Error(ErrorResponse::AudioUnavailable));
308                return;
309            }
310
311            let _ = response.send(AudioData { data }.into());
312        });
313    }
314
315    #[instrument(skip(self, response))]
316    fn handle_preload_check(
317        &self,
318        manifest_hash: Option<DataHash>,
319        response: oneshot::Sender<ServerMessage>,
320    ) {
321        if let Some(manifest_hash) = manifest_hash
322            && manifest_hash == self.audio_manifest_hash
323        {
324            let _ = response.send(PreloadMatch.into());
325            return;
326        }
327
328        let _ = response.send(
329            PreloadMismatch {
330                audio_manifest_hash: self.audio_manifest_hash.clone(),
331                achievements: self
332                    .achievements
333                    .values()
334                    .map(|achievement| AchievementRecord {
335                        id: achievement.id,
336                        audio_hash: achievement
337                            .audio_file
338                            .resolve(&self.audio_base_path)
339                            .as_ref()
340                            .map(|file| file.hash.clone()),
341                    })
342                    .collect(),
343            }
344            .into(),
345        );
346    }
347}
348
349#[cfg(feature = "tokio-graceful-shutdown")]
350#[derive(Debug, Error)]
351pub enum NeverError {}
352
353#[cfg(feature = "tokio-graceful-shutdown")]
354impl<Metadata, Player, State, Trigger> IntoSubsystem<NeverError>
355    for Engine<Metadata, Player, State, Trigger>
356where
357    Metadata: Send + Sync + 'static,
358    Player: PlayerInfo + PlayerMutator + Send + Sync + 'static,
359    State: Send + Sync + 'static,
360    Trigger: Copy + PartialEq + Eq + Debug + Send + Sync + 'static,
361{
362    async fn run(mut self, subsys: &mut SubsystemHandle) -> Result<(), NeverError> {
363        let _ = self.process_requests().cancel_on_shutdown(subsys).await;
364        Ok(())
365    }
366}
367
368#[derive(Debug, Error)]
369pub enum BuilderError {
370    #[error("missing field: {0}")]
371    MissingField(&'static str),
372}
373
374#[derive(Debug, Default)]
375pub struct EngineBuilder<Player, State = (), Trigger = (), Metadata = ()>
376where
377    Player: PlayerInfo + PlayerMutator,
378    Trigger: Copy,
379{
380    bloops: Vec<Bloop<Player>>,
381    achievements: Vec<Achievement<Metadata, Player, State, Trigger>>,
382    bloop_retention: Option<Duration>,
383    audio_base_path: Option<PathBuf>,
384    player_registry: Option<Arc<Mutex<PlayerRegistry<Player>>>>,
385    state: Option<Arc<Mutex<State>>>,
386    trigger_registry: Option<TriggerRegistry<Trigger>>,
387    network_rx: Option<mpsc::Receiver<(EngineRequest, oneshot::Sender<ServerMessage>)>>,
388    event_tx: Option<broadcast::Sender<Event>>,
389    throttle: Option<Throttle>,
390}
391
392impl<Player, State, Trigger, Metadata> EngineBuilder<Player, State, Trigger, Metadata>
393where
394    Player: PlayerInfo + PlayerMutator,
395    State: Default,
396    Trigger: Copy + PartialEq + Eq + Debug,
397{
398    pub fn new() -> Self {
399        Self {
400            bloops: Vec::new(),
401            achievements: Vec::new(),
402            bloop_retention: None,
403            audio_base_path: None,
404            player_registry: None,
405            state: None,
406            trigger_registry: None,
407            network_rx: None,
408            event_tx: None,
409            throttle: None,
410        }
411    }
412
413    pub fn bloops(mut self, bloops: Vec<Bloop<Player>>) -> Self {
414        self.bloops = bloops;
415        self
416    }
417
418    pub fn achievements(
419        mut self,
420        achievements: Vec<Achievement<Metadata, Player, State, Trigger>>,
421    ) -> Self {
422        self.achievements = achievements;
423        self
424    }
425
426    pub fn bloop_retention(mut self, retention: Duration) -> Self {
427        self.bloop_retention = Some(retention);
428        self
429    }
430
431    pub fn audio_base_path<P: Into<PathBuf>>(mut self, path: P) -> Self {
432        self.audio_base_path = Some(path.into());
433        self
434    }
435
436    pub fn player_registry(mut self, registry: Arc<Mutex<PlayerRegistry<Player>>>) -> Self {
437        self.player_registry = Some(registry);
438        self
439    }
440
441    pub fn state(mut self, state: Arc<Mutex<State>>) -> Self {
442        self.state = Some(state);
443        self
444    }
445
446    pub fn trigger_registry(mut self, registry: TriggerRegistry<Trigger>) -> Self {
447        self.trigger_registry = Some(registry);
448        self
449    }
450
451    pub fn network_rx(
452        mut self,
453        rx: mpsc::Receiver<(EngineRequest, oneshot::Sender<ServerMessage>)>,
454    ) -> Self {
455        self.network_rx = Some(rx);
456        self
457    }
458
459    pub fn event_tx(mut self, tx: broadcast::Sender<Event>) -> Self {
460        self.event_tx = Some(tx);
461        self
462    }
463
464    pub fn throttle(mut self, throttle: Throttle) -> Self {
465        self.throttle = Some(throttle);
466        self
467    }
468
469    /// Consumes the builder and constructs the Engine.
470    pub fn build(self) -> Result<Engine<Metadata, Player, State, Trigger>, BuilderError> {
471        let bloop_retention = self
472            .bloop_retention
473            .ok_or(BuilderError::MissingField("bloop_retention"))?;
474        let audio_base_path = self
475            .audio_base_path
476            .ok_or(BuilderError::MissingField("audio_base_path"))?;
477        let player_registry = self
478            .player_registry
479            .ok_or(BuilderError::MissingField("player_registry"))?;
480        let network_rx = self
481            .network_rx
482            .ok_or(BuilderError::MissingField("network_rx"))?;
483        let event_tx = self
484            .event_tx
485            .ok_or(BuilderError::MissingField("event_tx"))?;
486
487        let audio_manifest_hash = calculate_manifest_hash(&audio_base_path, &self.achievements);
488        let bloop_provider = BloopProvider::with_bloops(bloop_retention, self.bloops);
489        let achievements: HashMap<Uuid, Achievement<Metadata, Player, State, Trigger>> =
490            self.achievements.into_iter().map(|a| (a.id, a)).collect();
491        let state = self
492            .state
493            .unwrap_or_else(|| Arc::new(Mutex::new(Default::default())));
494        let trigger_registry = self
495            .trigger_registry
496            .unwrap_or_else(|| TriggerRegistry::new(HashMap::new()));
497        let hot_achievements = Vec::new();
498
499        Ok(Engine {
500            bloop_provider,
501            achievements,
502            audio_base_path,
503            audio_manifest_hash,
504            player_registry,
505            state,
506            trigger_registry,
507            hot_achievements,
508            network_rx,
509            event_tx,
510            throttle: self.throttle,
511        })
512    }
513}
514
515fn calculate_manifest_hash<Metadata, Player, State, Trigger>(
516    audio_base_path: &Path,
517    achievements: &[Achievement<Metadata, Player, State, Trigger>],
518) -> DataHash {
519    let audio_file_hashes: HashMap<Uuid, DataHash> = achievements
520        .iter()
521        .filter_map(|achievement| {
522            achievement
523                .audio_file
524                .resolve(audio_base_path)
525                .as_ref()
526                .map(|file| (achievement.id, file.hash.clone()))
527        })
528        .collect();
529
530    let mut entries: Vec<_> = audio_file_hashes.iter().collect();
531    entries.sort_by_key(|(id, _)| *id);
532    let mut hash_input = Vec::with_capacity(entries.len() * 32);
533
534    for (id, hash) in entries {
535        hash_input.extend(id.as_bytes());
536        hash_input.extend_from_slice(hash.as_bytes());
537    }
538
539    let manifest_hash = md5::compute(hash_input);
540    manifest_hash.into()
541}
542
543#[cfg(test)]
544mod tests {
545    use super::*;
546    use crate::test_utils::MockPlayer;
547    use crate::trigger::{TriggerOccurrence, TriggerSpec};
548
549    fn build_test_engine() -> Engine<(), MockPlayer, (), ()> {
550        let player_registry = Arc::new(Mutex::new(PlayerRegistry::new(vec![])));
551        let trigger_registry = TriggerRegistry::new(HashMap::new());
552
553        EngineBuilder::<MockPlayer, (), (), ()>::new()
554            .bloop_retention(Duration::from_secs(3600))
555            .audio_base_path("./audio")
556            .player_registry(player_registry)
557            .trigger_registry(trigger_registry)
558            .network_rx(mpsc::channel(1).1)
559            .event_tx(broadcast::channel(16).0)
560            .build()
561            .unwrap()
562    }
563
564    #[tokio::test]
565    async fn handle_bloop_rejects_unknown_nfc_uid() {
566        let mut engine = build_test_engine();
567        let unknown_nfc_uid = NfcUid::default();
568        let client_id = "test-client".to_string();
569
570        let (resp_tx, resp_rx) = oneshot::channel();
571        engine
572            .handle_bloop(unknown_nfc_uid, client_id.clone(), resp_tx)
573            .await;
574
575        let response = resp_rx.await.unwrap();
576        match response {
577            ServerMessage::Error(err) => {
578                assert!(matches!(err, ErrorResponse::UnknownNfcUid));
579            }
580            _ => panic!("Expected Error response for unknown NFC UID"),
581        }
582    }
583
584    #[tokio::test]
585    async fn handle_bloop_accepts_known_player() {
586        let mut engine = build_test_engine();
587        let nfc_uid = NfcUid::default();
588        let client_id = "test-client".to_string();
589
590        {
591            let mut registry = engine.player_registry.lock().await;
592            let (player, _) = MockPlayer::builder().nfc_uid(nfc_uid).build();
593            registry.add(Arc::into_inner(player).unwrap().into_inner().unwrap());
594        }
595
596        let (resp_tx, resp_rx) = oneshot::channel();
597        engine
598            .handle_bloop(nfc_uid, client_id.clone(), resp_tx)
599            .await;
600
601        let response = resp_rx.await.unwrap();
602        match response {
603            ServerMessage::BloopAccepted(BloopAccepted { achievements }) => {
604                assert!(achievements.is_empty());
605            }
606            _ => panic!("Expected BloopAccepted response"),
607        }
608    }
609
610    #[tokio::test]
611    async fn handle_bloop_activates_trigger_and_responds() {
612        let mut trigger_registry = HashMap::new();
613        trigger_registry.insert(
614            NfcUid::default(),
615            TriggerSpec {
616                trigger: (),
617                global: false,
618                occurrence: TriggerOccurrence::Once,
619            },
620        );
621        let trigger_registry = TriggerRegistry::new(trigger_registry);
622
623        let player_registry = Arc::new(Mutex::new(PlayerRegistry::new(vec![])));
624
625        let (_tx, rx) = mpsc::channel(1);
626        let (evt_tx, _) = broadcast::channel(16);
627
628        let mut engine = EngineBuilder::<MockPlayer, (), (), ()>::new()
629            .bloop_retention(Duration::from_secs(3600))
630            .audio_base_path("./audio")
631            .player_registry(player_registry)
632            .trigger_registry(trigger_registry)
633            .network_rx(rx)
634            .event_tx(evt_tx)
635            .build()
636            .unwrap();
637
638        let nfc_uid = NfcUid::default();
639        let client_id = "client".to_string();
640        let (resp_tx, resp_rx) = oneshot::channel();
641
642        engine
643            .handle_bloop(nfc_uid, client_id.clone(), resp_tx)
644            .await;
645
646        let response = resp_rx.await.unwrap();
647        match response {
648            ServerMessage::BloopAccepted(BloopAccepted { achievements }) => {
649                assert!(achievements.is_empty());
650            }
651            _ => panic!("Expected BloopAccepted response"),
652        }
653
654        assert!(
655            engine
656                .trigger_registry
657                .check_active_trigger((), "client", Utc::now())
658        );
659    }
660
661    #[tokio::test]
662    async fn handle_bloop_respects_throttling() {
663        let mut engine = build_test_engine();
664        let nfc_uid = NfcUid::default();
665        let client_id = "test-client".to_string();
666
667        {
668            let mut registry = engine.player_registry.lock().await;
669            let (player, _) = MockPlayer::builder().nfc_uid(nfc_uid).build();
670            registry.add(Arc::into_inner(player).unwrap().into_inner().unwrap());
671        }
672
673        engine.throttle = Some(Throttle::new(1, Duration::from_secs(10)));
674
675        let bloop = Bloop::new(
676            engine
677                .player_registry
678                .lock()
679                .await
680                .get_by_nfc_uid(nfc_uid)
681                .unwrap(),
682            client_id.clone(),
683            Utc::now(),
684        );
685        engine.bloop_provider.add(Arc::new(bloop));
686
687        let (resp_tx, resp_rx) = oneshot::channel();
688        engine
689            .handle_bloop(nfc_uid, client_id.clone(), resp_tx)
690            .await;
691
692        let response = resp_rx.await.unwrap();
693        match response {
694            ServerMessage::Error(err) => {
695                assert!(matches!(err, ErrorResponse::NfcUidThrottled));
696            }
697            _ => panic!("Expected throttling error"),
698        }
699    }
700}