Skip to main content

stoat/
client.rs

1use std::{panic::AssertUnwindSafe, sync::Arc, time::Duration};
2
3use futures::FutureExt;
4use stoat_database::events::client::EventV1;
5use tokio::{
6    select,
7    sync::{Mutex, mpsc},
8};
9
10use crate::{
11    Context, Error,
12    cache::GlobalCache,
13    context::Events,
14    events::{EventHandler, update_state},
15    http::HttpClient,
16    notifiers::Notifiers,
17    websocket::run,
18};
19
20#[derive(Clone)]
21pub struct Client<H> {
22    pub state: GlobalCache,
23    pub handler: Arc<H>,
24    pub http: HttpClient,
25    pub waiters: Notifiers,
26    events: Option<Events>,
27}
28
29impl<H: EventHandler + Clone + Send + Sync + 'static> Client<H> {
30    pub async fn new(handler: H) -> Result<Self, H::Error> {
31        Self::new_with_api_url(handler, "https://api.stoat.chat").await
32    }
33
34    pub async fn new_with_api_url(
35        handler: H,
36        base_url: impl Into<String>,
37    ) -> Result<Self, H::Error> {
38        let http = HttpClient::new(base_url.into(), None, None).await?;
39
40        Ok(Self {
41            state: GlobalCache::new((*http.api_config).clone()),
42            handler: Arc::new(handler),
43            http,
44            waiters: Notifiers::default(),
45            events: None,
46        })
47    }
48
49    pub async fn start(&mut self, token: impl Into<String>) -> Result<(), H::Error> {
50        let token = token.into();
51
52        self.http.token = Some(token.clone());
53        self.http.user_id = Some(self.http.fetch_self().await?.id);
54
55        Ok(())
56    }
57
58    pub async fn run(&mut self, token: impl Into<String>) -> Result<(), H::Error> {
59        let token = token.into();
60
61        self.start(token.clone()).await?;
62
63        let (client_sender, client_receiver) = mpsc::unbounded_channel();
64        self.events = Some(Events(Arc::new(client_sender)));
65
66        let (sender, receiver) = mpsc::unbounded_channel();
67
68        let handle = {
69            let sender = sender.clone();
70            let state = self.state.clone();
71            let token = token.clone();
72            let client_receiver = Arc::new(Mutex::new(client_receiver));
73
74            async move {
75                loop {
76                    if let Err(e) = run(
77                        sender.clone(),
78                        client_receiver.clone(),
79                        state.clone(),
80                        token.clone(),
81                    )
82                    .await
83                    {
84                        log::error!("{e:?}");
85
86                        if let Error::Close = e {
87                            return Ok(());
88                        }
89                    }
90
91                    log::info!("Disconnected! Reconnecting in 10 seconds.");
92
93                    tokio::time::sleep(Duration::from_secs(10)).await;
94                }
95            }
96        };
97
98        let res = select! {
99            e = handle => e,
100            _ = tokio::signal::ctrl_c() => {
101                log::info!("Received ctrl+c. exiting.");
102                Ok(())
103            }
104            _ = self.handle_events(receiver) => {
105                Ok(())
106            }
107        };
108
109        self.cleanup().await;
110
111        res
112    }
113
114    pub async fn cleanup(&mut self) {
115        self.state.cleanup().await;
116        self.waiters.clear_all_waiters().await;
117        self.events = None;
118    }
119
120    async fn handle_events(&self, mut receiver: mpsc::UnboundedReceiver<EventV1>) {
121        while let Some(event) = receiver.recv().await {
122            let this = self.clone();
123
124            tokio::spawn(async move {
125                this.handle_event(event).await;
126            });
127        }
128    }
129
130    pub async fn handle_event(&self, event: EventV1) {
131        let wrapper = AssertUnwindSafe(async {
132            let context = Context {
133                cache: self.state.clone(),
134                http: self.http.clone(),
135                notifiers: self.waiters.clone(),
136                events: self.events.clone().unwrap(),
137            };
138
139            update_state(event.clone(), context.clone(), self.handler.clone()).await
140        });
141
142        if let Err(e) = wrapper.catch_unwind().await {
143            log::error!("{e:?}");
144        }
145    }
146}