Skip to main content

lafere_api/
request_handlers.rs

1use std::{
2	any::{Any, TypeId},
3	collections::HashMap,
4	sync::Arc,
5};
6
7use lafere::{
8	handler::{Receiver, Sender},
9	packet::{Packet, PacketBytes},
10	server,
11};
12
13use crate::{
14	message::{Action, Message},
15	request::{EnableServerRequestsHandler, RequestHandler},
16	requestor::Requestor,
17	server::Session,
18};
19
20pub struct Data {
21	inner: HashMap<TypeId, Box<dyn Any + Send + Sync>>,
22}
23
24impl Data {
25	fn new() -> Self {
26		Self {
27			inner: HashMap::new(),
28		}
29	}
30
31	pub fn exists<D>(&self) -> bool
32	where
33		D: Any,
34	{
35		TypeId::of::<D>() == TypeId::of::<Session>()
36			|| self.inner.contains_key(&TypeId::of::<D>())
37	}
38
39	fn insert<D>(&mut self, data: D)
40	where
41		D: Any + Send + Sync,
42	{
43		self.inner.insert(data.type_id(), Box::new(data));
44	}
45
46	pub fn get<D>(&self) -> Option<&D>
47	where
48		D: Any,
49	{
50		self.inner
51			.get(&TypeId::of::<D>())
52			.and_then(|a| a.downcast_ref())
53	}
54
55	pub fn get_or_sess<'a, D>(&'a self, sess: &'a Session) -> Option<&'a D>
56	where
57		D: Any,
58	{
59		if TypeId::of::<D>() == TypeId::of::<Session>() {
60			<dyn Any>::downcast_ref(sess)
61		} else {
62			self.get()
63		}
64	}
65}
66
67struct Requests<A, B> {
68	inner: HashMap<A, Box<dyn RequestHandler<B, Action = A> + Send + Sync>>,
69	enable_server_requests: Option<
70		Box<dyn EnableServerRequestsHandler<B, Action = A> + Send + Sync>,
71	>,
72}
73
74impl<A, B> Requests<A, B>
75where
76	A: Action,
77{
78	fn new() -> Self {
79		Self {
80			inner: HashMap::new(),
81			enable_server_requests: None,
82		}
83	}
84
85	fn insert<H>(&mut self, handler: H)
86	where
87		H: RequestHandler<B, Action = A> + Send + Sync + 'static,
88	{
89		self.inner.insert(H::action(), Box::new(handler));
90	}
91
92	fn insert_enable_server_requests<H>(&mut self, handler: H)
93	where
94		H: EnableServerRequestsHandler<B, Action = A> + Send + Sync + 'static,
95	{
96		self.enable_server_requests = Some(Box::new(handler));
97	}
98
99	fn get(
100		&self,
101		action: &A,
102	) -> Option<&Box<dyn RequestHandler<B, Action = A> + Send + Sync>> {
103		self.inner.get(action)
104	}
105
106	fn get_enable_server_requests(
107		&self,
108	) -> Option<
109		&Box<dyn EnableServerRequestsHandler<B, Action = A> + Send + Sync>,
110	> {
111		self.enable_server_requests.as_ref()
112	}
113}
114
115pub struct RequestHandlersBuilder<A, B> {
116	requests: Requests<A, B>,
117	data: Data,
118}
119
120impl<A, B> RequestHandlersBuilder<A, B>
121where
122	A: Action,
123{
124	pub fn new() -> Self {
125		Self {
126			requests: Requests::new(),
127			data: Data::new(),
128		}
129	}
130
131	pub fn register_data<D>(&mut self, data: D)
132	where
133		D: Any + Send + Sync,
134	{
135		self.data.insert(data);
136	}
137
138	pub fn register_request<H>(&mut self, handler: H)
139	where
140		H: RequestHandler<B, Action = A> + Send + Sync + 'static,
141	{
142		handler.validate_data(&self.data);
143		self.requests.insert(handler);
144	}
145
146	pub(crate) fn register_enable_server_requests<H>(&mut self, handler: H)
147	where
148		H: EnableServerRequestsHandler<B, Action = A> + Send + Sync + 'static,
149	{
150		handler.validate_data(&self.data);
151		self.requests.insert_enable_server_requests(handler);
152	}
153
154	pub fn build(self) -> RequestHandlers<A, B> {
155		RequestHandlers(Arc::new(self))
156	}
157}
158
159pub struct RequestHandlers<A, B>(Arc<RequestHandlersBuilder<A, B>>);
160
161impl<A, B> RequestHandlers<A, B>
162where
163	A: Action,
164{
165	pub fn builder() -> RequestHandlersBuilder<A, B> {
166		RequestHandlersBuilder::new()
167	}
168
169	pub fn get_handler(
170		&self,
171		action: &A,
172	) -> Option<&Box<dyn RequestHandler<B, Action = A> + Send + Sync>> {
173		self.0.requests.get(action)
174	}
175
176	pub fn data(&self) -> &Data {
177		&self.0.data
178	}
179
180	pub fn get_data<D>(&self) -> Option<&D>
181	where
182		D: Any,
183	{
184		self.0.data.get()
185	}
186
187	pub(crate) async fn handle_connection(
188		&self,
189		session: Arc<Session>,
190		sender: Sender<Message<A, B>>,
191		mut recv: Receiver<Message<A, B>>,
192	) where
193		A: Send + Sync + 'static,
194		B: PacketBytes + Send + 'static,
195	{
196		while let Some(req) = recv.receive().await {
197			// todo replace with let else
198			let (msg, resp) = match req {
199				server::Message::Request(msg, resp) => (msg, resp),
200				server::Message::EnableServerRequests => {
201					self.enable_server_requests(
202						session.clone(),
203						sender.clone(),
204					);
205					continue;
206				}
207				// ignore streams for now
208				_ => continue,
209			};
210
211			let me = self.clone();
212			let session = session.clone();
213
214			let action = match msg.action() {
215				Some(act) => *act,
216				// todo once we bump the version again
217				// we need to pass our own errors via packets
218				// not only those from the api users
219				None => {
220					tracing::error!("invalid action received");
221					continue;
222				}
223			};
224
225			tokio::spawn(async move {
226				let handler = match me.get_handler(&action) {
227					Some(handler) => handler,
228					// todo once we bump the version again
229					// we need to pass our own errors via packets
230					// not only those from the api users
231					None => {
232						tracing::error!("no handler for {:?}", action);
233						return;
234					}
235				};
236				let r = handler.handle(msg, &me.data(), &session).await;
237
238				match r {
239					Ok(mut msg) => {
240						msg.header_mut().set_action(action);
241						// i don't care about the response
242						let _ = resp.send(msg);
243					}
244					Err(e) => {
245						// todo once we bump the version again
246						// we need to pass our own errors via packets
247						// not only those from the api users
248						tracing::error!("handler returned an error {:?}", e);
249					}
250				}
251			});
252		}
253	}
254
255	fn enable_server_requests(
256		&self,
257		session: Arc<Session>,
258		sender: Sender<Message<A, B>>,
259	) where
260		A: Send + Sync + 'static,
261		B: PacketBytes + Send + 'static,
262	{
263		let me = self.clone();
264
265		tokio::spawn(async move {
266			let handler = match me.0.requests.get_enable_server_requests() {
267				Some(handler) => handler,
268				// todo once we bump the version again
269				// we need to pass our own errors via packets
270				// not only those from the api users
271				None => {
272					tracing::error!(
273						"no handler for enable server requests found"
274					);
275					return;
276				}
277			};
278
279			let r = handler
280				.handle(Requestor::new(sender), me.data(), &session)
281				.await;
282			if let Err(e) = r {
283				// todo once we bump the version again
284				// we need to pass our own errors via packets
285				// not only those from the api users
286				tracing::error!("handler returned an error {:?}", e);
287			}
288		});
289	}
290}
291
292impl<A, B> Clone for RequestHandlers<A, B> {
293	fn clone(&self) -> Self {
294		Self(Arc::clone(&self.0))
295	}
296}