1use std::future::Future;
2use std::io::Cursor;
3use std::marker::PhantomData;
4use std::pin::Pin;
5use std::sync::Arc;
6
7use anyhow::{Context as _, Result};
8use rivetkit_core::{
9 ActorContext, EnqueueAndWaitOpts, QueueMessage as CoreQueueMessage, QueueNextBatchOpts,
10 QueueNextOpts, QueueTryNextBatchOpts, QueueTryNextOpts,
11};
12use serde::{Serialize, de::DeserializeOwned};
13
14use crate::{actor::Actor, context::Ctx};
15pub(crate) type BoxQueueFuture = Pin<Box<dyn Future<Output = Result<Option<Vec<u8>>>> + Send>>;
16
17pub trait QueueMessage: Serialize + DeserializeOwned + Send + Sync + 'static {
18 type Reply: Serialize + DeserializeOwned + Send + 'static;
19
20 const NAME: &'static str;
21}
22
23pub trait HandlesQueue<M: QueueMessage>: Actor + Sized {
24 type Future: Future<Output = Result<M::Reply>> + Send + 'static;
25
26 fn handle_queue(self: Arc<Self>, ctx: Ctx<Self>, message: M) -> Self::Future;
27}
28
29#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30pub struct QueueEntry<A: Actor> {
31 pub name: &'static str,
32 _p: PhantomData<fn() -> A>,
33}
34
35impl<A: Actor> QueueEntry<A> {
36 pub const fn new(name: &'static str) -> Self {
37 Self {
38 name,
39 _p: PhantomData,
40 }
41 }
42}
43
44pub trait QueueSet<A: Actor>: Send + Sync + 'static {
45 fn entries() -> Vec<QueueEntry<A>>;
46 fn dispatch(actor: Arc<A>, ctx: Ctx<A>, name: &str, body: &[u8]) -> Option<BoxQueueFuture>;
47}
48
49impl<A: Actor> QueueSet<A> for () {
50 fn entries() -> Vec<QueueEntry<A>> {
51 Vec::new()
52 }
53
54 fn dispatch(_actor: Arc<A>, _ctx: Ctx<A>, _name: &str, _body: &[u8]) -> Option<BoxQueueFuture> {
55 None
56 }
57}
58
59macro_rules! impl_queue_set {
60 ($($message:ident),+) => {
61 impl<Act, $($message),+> QueueSet<Act> for ($($message,)+)
62 where
63 Act: Actor + $(HandlesQueue<$message> +)+,
64 $($message: QueueMessage,)+
65 {
66 fn entries() -> Vec<QueueEntry<Act>> {
67 vec![$(QueueEntry::new(<$message as QueueMessage>::NAME)),+]
68 }
69
70 fn dispatch(
71 actor: Arc<Act>,
72 ctx: Ctx<Act>,
73 name: &str,
74 body: &[u8],
75 ) -> Option<BoxQueueFuture> {
76 $(
77 if name == <$message as QueueMessage>::NAME {
78 let body = body.to_vec();
79 return Some(Box::pin(async move {
80 let message = decode_cbor::<$message>(&body, "queue message body")
81 .with_context(|| format!("decode queue message '{}'", <$message as QueueMessage>::NAME))?;
82 let reply = <Act as HandlesQueue<$message>>::handle_queue(actor, ctx, message).await?;
83 Ok(Some(encode_cbor(&reply, "queue message reply")?))
84 }));
85 }
86 )+
87 None
88 }
89 }
90 };
91}
92
93impl_queue_set!(M0);
94impl_queue_set!(M0, M1);
95impl_queue_set!(M0, M1, M2);
96impl_queue_set!(M0, M1, M2, M3);
97impl_queue_set!(M0, M1, M2, M3, M4);
98impl_queue_set!(M0, M1, M2, M3, M4, M5);
99impl_queue_set!(M0, M1, M2, M3, M4, M5, M6);
100impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7);
101impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8);
102impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8, M9);
103impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10);
104impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10, M11);
105impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10, M11, M12);
106impl_queue_set!(M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10, M11, M12, M13);
107impl_queue_set!(
108 M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10, M11, M12, M13, M14
109);
110impl_queue_set!(
111 M0, M1, M2, M3, M4, M5, M6, M7, M8, M9, M10, M11, M12, M13, M14, M15
112);
113
114pub struct Queue<'a, A: Actor> {
122 inner: &'a ActorContext,
123 _p: PhantomData<fn() -> A>,
124}
125
126#[derive(Debug, Clone)]
127pub struct TypedQueueMessage<M: QueueMessage> {
128 inner: CoreQueueMessage,
129 body: M,
130}
131
132impl<M: QueueMessage> TypedQueueMessage<M> {
133 pub fn id(&self) -> u64 {
134 self.inner.id
135 }
136
137 pub fn name(&self) -> &str {
138 &self.inner.name
139 }
140
141 pub fn body(&self) -> &M {
142 &self.body
143 }
144
145 pub fn created_at(&self) -> i64 {
146 self.inner.created_at
147 }
148
149 pub fn is_completable(&self) -> bool {
150 self.inner.is_completable()
151 }
152
153 pub fn into_body(self) -> M {
154 self.body
155 }
156
157 pub fn as_core(&self) -> &CoreQueueMessage {
158 &self.inner
159 }
160
161 pub fn into_core(self) -> CoreQueueMessage {
162 self.inner
163 }
164
165 pub async fn complete(self, reply: M::Reply) -> Result<()> {
166 self.complete_raw(Some(encode_cbor(&reply, "queue message reply")?))
167 .await
168 }
169
170 pub async fn complete_raw(self, response: Option<Vec<u8>>) -> Result<()> {
171 self.inner.complete(response).await
172 }
173}
174
175impl<'a, A: Actor> Queue<'a, A> {
176 pub(crate) fn new(inner: &'a ActorContext) -> Self {
177 Self {
178 inner,
179 _p: PhantomData,
180 }
181 }
182
183 pub async fn send<T: Serialize>(&self, name: &str, body: &T) -> Result<CoreQueueMessage> {
185 self.send_raw(name, &encode_cbor(body, "queue message body")?)
186 .await
187 }
188
189 pub async fn send_raw(&self, name: &str, body: &[u8]) -> Result<CoreQueueMessage> {
191 self.inner.send(name, body).await
192 }
193
194 pub async fn enqueue_and_wait<T: Serialize>(
197 &self,
198 name: &str,
199 body: &T,
200 opts: EnqueueAndWaitOpts,
201 ) -> Result<Option<Vec<u8>>> {
202 self.enqueue_and_wait_raw(name, &encode_cbor(body, "queue message body")?, opts)
203 .await
204 }
205
206 pub async fn enqueue_and_wait_raw(
208 &self,
209 name: &str,
210 body: &[u8],
211 opts: EnqueueAndWaitOpts,
212 ) -> Result<Option<Vec<u8>>> {
213 self.inner.enqueue_and_wait(name, body, opts).await
214 }
215
216 pub async fn next(&self, opts: QueueNextOpts) -> Result<Option<CoreQueueMessage>> {
218 self.inner.next(opts).await
219 }
220
221 pub async fn next_typed<M: QueueMessage>(
223 &self,
224 opts: QueueNextOpts,
225 ) -> Result<Option<TypedQueueMessage<M>>> {
226 self.inner
227 .next(typed_next_opts::<M>(opts))
228 .await?
229 .map(decode_core_message)
230 .transpose()
231 }
232
233 pub async fn next_batch(&self, opts: QueueNextBatchOpts) -> Result<Vec<CoreQueueMessage>> {
235 self.inner.next_batch(opts).await
236 }
237
238 pub async fn next_batch_typed<M: QueueMessage>(
240 &self,
241 opts: QueueNextBatchOpts,
242 ) -> Result<Vec<TypedQueueMessage<M>>> {
243 self.inner
244 .next_batch(typed_next_batch_opts::<M>(opts))
245 .await?
246 .into_iter()
247 .map(decode_core_message)
248 .collect()
249 }
250
251 pub fn try_next(&self, opts: QueueTryNextOpts) -> Result<Option<CoreQueueMessage>> {
253 self.inner.try_next(opts)
254 }
255
256 pub fn try_next_typed<M: QueueMessage>(
258 &self,
259 opts: QueueTryNextOpts,
260 ) -> Result<Option<TypedQueueMessage<M>>> {
261 self.inner
262 .try_next(typed_try_next_opts::<M>(opts))?
263 .map(decode_core_message)
264 .transpose()
265 }
266
267 pub fn try_next_batch(&self, opts: QueueTryNextBatchOpts) -> Result<Vec<CoreQueueMessage>> {
269 self.inner.try_next_batch(opts)
270 }
271
272 pub fn try_next_batch_typed<M: QueueMessage>(
274 &self,
275 opts: QueueTryNextBatchOpts,
276 ) -> Result<Vec<TypedQueueMessage<M>>> {
277 self.inner
278 .try_next_batch(typed_try_next_batch_opts::<M>(opts))?
279 .into_iter()
280 .map(decode_core_message)
281 .collect()
282 }
283
284 pub async fn inspect_messages(&self) -> Result<Vec<CoreQueueMessage>> {
286 self.inner.inspect_messages().await
287 }
288
289 pub fn max_size(&self) -> u32 {
291 self.inner.max_size()
292 }
293}
294
295fn decode_core_message<M: QueueMessage>(message: CoreQueueMessage) -> Result<TypedQueueMessage<M>> {
296 if message.name != M::NAME {
297 anyhow::bail!(
298 "expected queue message '{}', received '{}'",
299 M::NAME,
300 message.name
301 );
302 }
303
304 let body = decode_cbor::<M>(&message.body, "queue message body")
305 .with_context(|| format!("decode queue message '{}'", M::NAME))?;
306
307 Ok(TypedQueueMessage {
308 inner: message,
309 body,
310 })
311}
312
313fn typed_next_opts<M: QueueMessage>(opts: QueueNextOpts) -> QueueNextOpts {
314 QueueNextOpts {
315 names: Some(vec![M::NAME.to_owned()]),
316 timeout: opts.timeout,
317 signal: opts.signal,
318 completable: opts.completable,
319 }
320}
321
322fn typed_next_batch_opts<M: QueueMessage>(opts: QueueNextBatchOpts) -> QueueNextBatchOpts {
323 QueueNextBatchOpts {
324 names: Some(vec![M::NAME.to_owned()]),
325 count: opts.count,
326 timeout: opts.timeout,
327 signal: opts.signal,
328 completable: opts.completable,
329 }
330}
331
332fn typed_try_next_opts<M: QueueMessage>(opts: QueueTryNextOpts) -> QueueTryNextOpts {
333 QueueTryNextOpts {
334 names: Some(vec![M::NAME.to_owned()]),
335 completable: opts.completable,
336 }
337}
338
339fn typed_try_next_batch_opts<M: QueueMessage>(
340 opts: QueueTryNextBatchOpts,
341) -> QueueTryNextBatchOpts {
342 QueueTryNextBatchOpts {
343 names: Some(vec![M::NAME.to_owned()]),
344 count: opts.count,
345 completable: opts.completable,
346 }
347}
348
349fn encode_cbor<T: Serialize>(value: &T, label: &str) -> Result<Vec<u8>> {
350 let mut encoded = Vec::new();
351 ciborium::into_writer(value, &mut encoded)
352 .with_context(|| format!("encode {label} as cbor"))?;
353 Ok(encoded)
354}
355
356fn decode_cbor<T: DeserializeOwned>(bytes: &[u8], label: &str) -> Result<T> {
357 ciborium::from_reader(Cursor::new(bytes)).with_context(|| format!("decode {label} from cbor"))
358}
359
360#[cfg(test)]
361mod tests {
362 use std::future::{Ready, ready};
363 use std::sync::Arc;
364
365 use anyhow::Result;
366 use serde::{Deserialize, Serialize};
367
368 use super::{HandlesQueue, QueueMessage, QueueSet};
369 use crate::{action, actor::Actor, context::Ctx};
370
371 const QUEUE_SET_TUPLE_ARITY_MAX: usize = 16;
372
373 struct TestActor;
374
375 impl Actor for TestActor {
376 type State = ();
377 type Input = ();
378 type Actions = ();
379 type Events = ();
380 type Queue = ();
381 type ConnParams = ();
382 type ConnState = ();
383 type Action = action::Raw;
384 }
385
386 #[derive(Debug, Serialize, Deserialize)]
387 struct FirstMessage;
388
389 impl QueueMessage for FirstMessage {
390 type Reply = ();
391
392 const NAME: &'static str = "first";
393 }
394
395 #[derive(Debug, Serialize, Deserialize)]
396 struct SecondMessage;
397
398 impl QueueMessage for SecondMessage {
399 type Reply = ();
400
401 const NAME: &'static str = "second";
402 }
403
404 impl HandlesQueue<FirstMessage> for TestActor {
405 type Future = Ready<Result<()>>;
406
407 fn handle_queue(self: Arc<Self>, _ctx: Ctx<Self>, _message: FirstMessage) -> Self::Future {
408 ready(Ok(()))
409 }
410 }
411
412 impl HandlesQueue<SecondMessage> for TestActor {
413 type Future = Ready<Result<()>>;
414
415 fn handle_queue(self: Arc<Self>, _ctx: Ctx<Self>, _message: SecondMessage) -> Self::Future {
416 ready(Ok(()))
417 }
418 }
419
420 #[test]
421 fn queue_set_unit_registers_nothing() {
422 assert!(<() as QueueSet<TestActor>>::entries().is_empty());
423 }
424
425 #[test]
426 fn queue_set_tuple_registers_names_in_order() {
427 let entries = <(FirstMessage, SecondMessage) as QueueSet<TestActor>>::entries();
428
429 assert_eq!(
430 entries.iter().map(|entry| entry.name).collect::<Vec<_>>(),
431 ["first", "second",]
432 );
433 }
434
435 #[test]
436 fn queue_set_tuple_supports_one_and_max_arity() {
437 assert_eq!(
438 <(FirstMessage,) as QueueSet<TestActor>>::entries()
439 .iter()
440 .map(|entry| entry.name)
441 .collect::<Vec<_>>(),
442 ["first"]
443 );
444
445 type MaxMessages = (
446 FirstMessage,
447 FirstMessage,
448 FirstMessage,
449 FirstMessage,
450 FirstMessage,
451 FirstMessage,
452 FirstMessage,
453 FirstMessage,
454 FirstMessage,
455 FirstMessage,
456 FirstMessage,
457 FirstMessage,
458 FirstMessage,
459 FirstMessage,
460 FirstMessage,
461 FirstMessage,
462 );
463 let entries = <MaxMessages as QueueSet<TestActor>>::entries();
464
465 assert_eq!(entries.len(), QUEUE_SET_TUPLE_ARITY_MAX);
466 assert!(entries.iter().all(|entry| entry.name == "first"));
467 }
468}