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}