Skip to main content

unb_runtime/
client.rs

1use std::collections::HashMap;
2use std::sync::{Arc, Mutex};
3
4use unb_core::{
5    ClientDelivery as CoreClientDelivery, ClientOperationId, DiscoverPlan, Envelope, ErrorCode,
6    Kind,
7};
8use bytes::Bytes;
9use serde_json::Value;
10use tokio::sync::mpsc;
11
12use crate::wire::Directive;
13use crate::{BodyStream, WireBody};
14
15#[derive(Debug, thiserror::Error, PartialEq, Eq)]
16pub enum ClientError {
17    #[error("client operation {0} is gone")]
18    Gone(String),
19    #[error("client operation {operation:?} failed with {code:?}: {message}")]
20    Protocol {
21        operation: ClientOperationId,
22        code: ErrorCode,
23        message: String,
24    },
25    #[error("client operation {0} was cancelled")]
26    Cancelled(String),
27    #[error("client operation {0} timed out")]
28    Timeout(String),
29    #[error("client session closed during operation {0}")]
30    SessionClosed(String),
31    #[error("invalid client request: {0}")]
32    Invalid(String),
33}
34
35const CLIENT_OPERATION_QUEUE: usize = 64;
36
37#[derive(Clone)]
38pub struct ClientSession {
39    operations: Arc<Mutex<HashMap<ClientOperationId, mpsc::Sender<ClientDelivery>>>>,
40    directives: mpsc::Sender<Directive>,
41}
42
43impl ClientSession {
44    pub(crate) fn connected(directives: mpsc::Sender<Directive>) -> Self {
45        Self {
46            operations: Arc::default(),
47            directives,
48        }
49    }
50
51    pub async fn start(
52        &self,
53        subject: &str,
54        kind: Kind,
55        payload: Bytes,
56        hops: Option<u8>,
57        headers: serde_json::Map<String, Value>,
58    ) -> Result<ClientStream, ClientError> {
59        self.start_with_timeout(subject, kind, payload, hops, headers, None, None)
60            .await
61    }
62
63    pub async fn start_streaming(
64        &self,
65        subject: &str,
66        kind: Kind,
67        body: BodyStream,
68        hops: Option<u8>,
69        headers: serde_json::Map<String, Value>,
70    ) -> Result<ClientStream, ClientError> {
71        self.start_with_timeout(subject, kind, Bytes::new(), hops, headers, Some(body), None)
72            .await
73    }
74
75    pub async fn fetch(
76        &self,
77        request: http::Request<Bytes>,
78        timeout: std::time::Duration,
79    ) -> Result<http::Response<Bytes>, ClientError> {
80        let (parts, body) = request.into_parts();
81        let response = self
82            .fetch_body(
83                http::Request::from_parts(parts, WireBody::Bytes(body)),
84                timeout,
85            )
86            .await?;
87        let (parts, body) = response.into_parts();
88        let body = body
89            .collect_to(unb_transport::DEFAULT_MAX_FRAME_SIZE)
90            .await
91            .map_err(|error| ClientError::Invalid(error.to_string()))?;
92        Ok(http::Response::from_parts(parts, body))
93    }
94
95    pub async fn fetch_body(
96        &self,
97        request: http::Request<WireBody>,
98        timeout: std::time::Duration,
99    ) -> Result<http::Response<WireBody>, ClientError> {
100        let (parts, body) = request.into_parts();
101        let envelope = Envelope::from_request(http::Request::from_parts(parts, Bytes::new()))
102            .map_err(|error| ClientError::Invalid(error.to_string()))?;
103        let (payload, body) = match body {
104            WireBody::Bytes(payload) => (payload, None),
105            WireBody::Stream(body) => (Bytes::new(), Some(body)),
106        };
107        let mut stream = self
108            .start_with_timeout(
109                &envelope.subject,
110                Kind::Request,
111                payload,
112                envelope.hops,
113                envelope.headers,
114                body,
115                Some(timeout),
116            )
117            .await?;
118        let response = stream
119            .next_response()
120            .await?
121            .ok_or_else(|| ClientError::Gone(stream.operation().as_str().to_owned()))?;
122        response
123            .into_response()
124            .map_err(|error| ClientError::Invalid(error.to_string()))
125    }
126
127    pub async fn subscribe(
128        &self,
129        request: http::Request<Bytes>,
130        timeout: Option<std::time::Duration>,
131    ) -> Result<ClientStream, ClientError> {
132        let envelope = Envelope::from_request(request)
133            .map_err(|error| ClientError::Invalid(error.to_string()))?;
134        self.start_with_timeout(
135            &envelope.subject,
136            Kind::Subscribe,
137            envelope.payload,
138            envelope.hops,
139            envelope.headers,
140            None,
141            timeout,
142        )
143        .await
144    }
145
146    pub async fn discover(&self, plan: DiscoverPlan) -> Result<ClientStream, ClientError> {
147        let timeout = plan.timeout_ms.map(std::time::Duration::from_millis);
148        let hops = Some(plan.hops);
149        let payload = Envelope::encode_payload(
150            &serde_json::to_value(plan).map_err(|error| ClientError::Invalid(error.to_string()))?,
151        );
152        self.start_with_timeout(
153            "",
154            Kind::Discover,
155            payload,
156            hops,
157            Default::default(),
158            None,
159            timeout,
160        )
161        .await
162    }
163
164    #[allow(clippy::too_many_arguments)]
165    async fn start_with_timeout(
166        &self,
167        subject: &str,
168        kind: Kind,
169        payload: Bytes,
170        hops: Option<u8>,
171        headers: serde_json::Map<String, Value>,
172        body: Option<BodyStream>,
173        timeout: Option<std::time::Duration>,
174    ) -> Result<ClientStream, ClientError> {
175        let (sender, receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
176        let (reply, response) = tokio::sync::oneshot::channel();
177        self.directives
178            .send(Directive::StartClientOperation {
179                subject: subject.to_owned(),
180                kind,
181                payload,
182                hops,
183                headers,
184                body,
185                timeout,
186                sender,
187                reply,
188            })
189            .await
190            .map_err(|_| ClientError::Gone("session".into()))?;
191        let operation = response
192            .await
193            .map_err(|_| ClientError::Gone("session".into()))?
194            .map_err(|error| ClientError::Invalid(error.to_string()))?;
195        Ok(ClientStream {
196            operation,
197            receiver,
198            directives: self.directives.clone(),
199            completed: false,
200            pending: None,
201        })
202    }
203
204    pub(crate) fn register(
205        &self,
206        operation: ClientOperationId,
207        sender: mpsc::Sender<ClientDelivery>,
208    ) -> Result<(), ClientError> {
209        let mut operations = self.operations.lock().expect("client operations lock");
210        if operations.insert(operation.clone(), sender).is_some() {
211            return Err(ClientError::Gone(operation.as_str().to_owned()));
212        }
213        Ok(())
214    }
215
216    pub(crate) async fn deliver(
217        &self,
218        operation: &ClientOperationId,
219        delivery: CoreClientDelivery,
220        body: Option<WireBody>,
221    ) {
222        let terminal = matches!(
223            delivery,
224            CoreClientDelivery::Terminal(_)
225                | CoreClientDelivery::Cancelled
226                | CoreClientDelivery::TimedOut
227                | CoreClientDelivery::SessionClosed
228        );
229        let sender = self
230            .operations
231            .lock()
232            .expect("client operations lock")
233            .get(operation)
234            .cloned();
235        let Some(sender) = sender else { return };
236        let delivered = sender
237            .send(ClientDelivery {
238                core: delivery,
239                body,
240            })
241            .await
242            .is_ok();
243        if terminal || !delivered {
244            self.abandon(operation);
245        }
246    }
247
248    pub(crate) fn abandon(&self, operation: &ClientOperationId) {
249        self.operations
250            .lock()
251            .expect("client operations lock")
252            .remove(operation);
253    }
254}
255
256#[doc(hidden)]
257pub struct ClientDelivery {
258    core: CoreClientDelivery,
259    body: Option<WireBody>,
260}
261
262pub struct ClientStream {
263    operation: ClientOperationId,
264    receiver: mpsc::Receiver<ClientDelivery>,
265    directives: mpsc::Sender<Directive>,
266    completed: bool,
267    pending: Option<ClientDelivery>,
268}
269
270impl ClientStream {
271    pub fn operation(&self) -> &ClientOperationId {
272        &self.operation
273    }
274
275    pub async fn next(&mut self) -> Result<Option<Envelope>, ClientError> {
276        if self.completed {
277            return Ok(None);
278        }
279        let delivery = match self.pending.take() {
280            Some(delivery) => Some(delivery),
281            None => self.receiver.recv().await,
282        };
283        match delivery {
284            Some(delivery) => {
285                let response = self.project_head(delivery)?;
286                match response {
287                    Some(response) => {
288                        let (mut envelope, body) = response.into_parts();
289                        envelope.payload = body
290                            .collect_to(unb_transport::DEFAULT_MAX_FRAME_SIZE)
291                            .await
292                            .map_err(|error| ClientError::Invalid(error.to_string()))?;
293                        Ok(Some(envelope))
294                    }
295                    None => Ok(None),
296                }
297            }
298            None => {
299                self.completed = true;
300                Err(ClientError::Gone(self.operation.as_str().to_owned()))
301            }
302        }
303    }
304
305    pub fn try_next(&mut self) -> Result<Option<Option<Envelope>>, ClientError> {
306        if self.completed {
307            return Ok(Some(None));
308        }
309        match self.receiver.try_recv() {
310            Ok(mut delivery) => match delivery.body.take() {
311                Some(WireBody::Stream(stream)) => {
312                    delivery.body = Some(WireBody::Stream(stream));
313                    self.pending = Some(delivery);
314                    Ok(None)
315                }
316                Some(WireBody::Bytes(payload)) => {
317                    let mut envelope = self.project_without_body(delivery)?;
318                    if let Some(envelope) = &mut envelope {
319                        envelope.payload = payload;
320                    }
321                    Ok(Some(envelope))
322                }
323                None => self.project_without_body(delivery).map(Some),
324            },
325            Err(mpsc::error::TryRecvError::Empty) => Ok(None),
326            Err(mpsc::error::TryRecvError::Disconnected) => {
327                self.completed = true;
328                Err(ClientError::Gone(self.operation.as_str().to_owned()))
329            }
330        }
331    }
332
333    pub async fn next_response(&mut self) -> Result<Option<ClientResponse>, ClientError> {
334        if self.completed {
335            return Ok(None);
336        }
337        let delivery = match self.pending.take() {
338            Some(delivery) => Some(delivery),
339            None => self.receiver.recv().await,
340        };
341        match delivery {
342            Some(delivery) => self.project_head(delivery),
343            None => {
344                self.completed = true;
345                Err(ClientError::Gone(self.operation.as_str().to_owned()))
346            }
347        }
348    }
349
350    fn project_without_body(
351        &mut self,
352        delivery: ClientDelivery,
353    ) -> Result<Option<Envelope>, ClientError> {
354        self.project_head(delivery)
355            .map(|response| response.map(|response| response.envelope))
356    }
357
358    fn project_head(
359        &mut self,
360        delivery: ClientDelivery,
361    ) -> Result<Option<ClientResponse>, ClientError> {
362        match delivery.core {
363            CoreClientDelivery::Terminal(frame) if frame.head.kind == Kind::Error => {
364                self.completed = true;
365                let error = frame.head.error.unwrap_or(unb_core::ApplicationError {
366                    code: ErrorCode::Protocol,
367                    message: "protocol error".to_owned(),
368                });
369                Err(ClientError::Protocol {
370                    operation: self.operation.clone(),
371                    code: error.code,
372                    message: error.message,
373                })
374            }
375            CoreClientDelivery::Terminal(frame) => {
376                self.completed = true;
377                Ok(Some(ClientResponse {
378                    envelope: frame.into_envelope(),
379                    body: delivery
380                        .body
381                        .unwrap_or_else(|| WireBody::Bytes(Bytes::new())),
382                }))
383            }
384            CoreClientDelivery::Item(frame) => Ok(Some(ClientResponse {
385                envelope: frame.into_envelope(),
386                body: delivery
387                    .body
388                    .unwrap_or_else(|| WireBody::Bytes(Bytes::new())),
389            })),
390            CoreClientDelivery::Cancelled => {
391                self.completed = true;
392                Err(ClientError::Cancelled(self.operation.as_str().to_owned()))
393            }
394            CoreClientDelivery::TimedOut => {
395                self.completed = true;
396                Err(ClientError::Timeout(self.operation.as_str().to_owned()))
397            }
398            CoreClientDelivery::SessionClosed => {
399                self.completed = true;
400                Err(ClientError::SessionClosed(
401                    self.operation.as_str().to_owned(),
402                ))
403            }
404        }
405    }
406}
407
408pub struct ClientResponse {
409    envelope: Envelope,
410    body: WireBody,
411}
412
413impl ClientResponse {
414    pub fn head(&self) -> &Envelope {
415        &self.envelope
416    }
417
418    pub fn into_parts(self) -> (Envelope, WireBody) {
419        (self.envelope, self.body)
420    }
421
422    pub fn into_body(self) -> WireBody {
423        self.body
424    }
425
426    pub fn into_response(self) -> Result<http::Response<WireBody>, unb_core::CoreError> {
427        let response = self.envelope.to_response()?;
428        let (parts, _) = response.into_parts();
429        Ok(http::Response::from_parts(parts, self.body))
430    }
431}
432
433impl Drop for ClientStream {
434    fn drop(&mut self) {
435        if self.completed {
436            return;
437        }
438        self.completed = true;
439        let command = Directive::CancelClientOperation {
440            operation: self.operation.clone(),
441        };
442        if let Err(mpsc::error::TrySendError::Full(command)) = self.directives.try_send(command) {
443            let directives = self.directives.clone();
444            n0_future::task::spawn(async move {
445                let _ = directives.send(command).await;
446            });
447        }
448    }
449}
450
451#[cfg(test)]
452mod tests {
453    use super::*;
454
455    fn event(sequence: usize) -> Envelope {
456        Envelope {
457            v: 1,
458            id: sequence.to_string(),
459            subject: String::new(),
460            kind: Kind::Event,
461            corr: None,
462            seq: None,
463            hops: None,
464            body_token: None,
465            payload: Bytes::new(),
466            path: Vec::new(),
467            headers: Default::default(),
468        }
469    }
470
471    #[tokio::test]
472    async fn saturated_delivery_blocks_without_loss_until_capacity_returns() {
473        let (directives, _directive_receiver) = mpsc::channel(1);
474        let session = ClientSession::connected(directives);
475        let operation = ClientOperationId::from("full");
476        let (sender, mut receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
477        session.register(operation.clone(), sender).unwrap();
478
479        let total = CLIENT_OPERATION_QUEUE * 4;
480        let producer = tokio::spawn({
481            let session = session.clone();
482            let operation = operation.clone();
483            async move {
484                for sequence in 0..total {
485                    session
486                        .deliver(
487                            &operation,
488                            CoreClientDelivery::Item(
489                                unb_core::ApplicationFrame::from_envelope(&event(sequence))
490                                    .unwrap(),
491                            ),
492                            None,
493                        )
494                        .await;
495                }
496                let terminal = Envelope {
497                    v: 1,
498                    id: "terminal".into(),
499                    subject: String::new(),
500                    kind: Kind::Error,
501                    corr: Some(operation.as_str().to_owned()),
502                    seq: None,
503                    hops: None,
504                    body_token: None,
505                    payload: Envelope::encode_payload(&serde_json::json!({
506                        "code": ErrorCode::Busy,
507                        "message": "busy"
508                    })),
509                    path: Vec::new(),
510                    headers: Default::default(),
511                };
512                session
513                    .deliver(
514                        &operation,
515                        CoreClientDelivery::Terminal(
516                            unb_core::ApplicationFrame::from_envelope(&terminal).unwrap(),
517                        ),
518                        None,
519                    )
520                    .await;
521            }
522        });
523
524        for sequence in 0..total {
525            match receiver.recv().await.unwrap().core {
526                CoreClientDelivery::Item(frame) => {
527                    assert_eq!(frame.head.id, sequence.to_string())
528                }
529                _ => panic!("expected ordered item {sequence}"),
530            }
531        }
532        assert!(matches!(
533            receiver.recv().await,
534            Some(ClientDelivery {
535                core: CoreClientDelivery::Terminal(frame),
536                ..
537            }) if frame.head.error.as_ref().is_some_and(|error| error.code == ErrorCode::Busy)
538        ));
539        producer.await.unwrap();
540        assert!(!session
541            .operations
542            .lock()
543            .expect("client operations lock")
544            .contains_key(&operation));
545    }
546
547    #[tokio::test]
548    async fn response_head_is_delivered_before_its_body_completes() {
549        let (directives, _directive_receiver) = mpsc::channel(1);
550        let (sender, receiver) = mpsc::channel(1);
551        let mut stream = ClientStream {
552            operation: ClientOperationId::from("head-first"),
553            receiver,
554            directives,
555            completed: false,
556            pending: None,
557        };
558        let response = Envelope {
559            v: 1,
560            id: "response".into(),
561            subject: String::new(),
562            kind: Kind::Response,
563            corr: Some("head-first".into()),
564            seq: None,
565            hops: None,
566            body_token: Some("body-1".into()),
567            payload: Bytes::new(),
568            path: Vec::new(),
569            headers: serde_json::Map::from_iter([(
570                "x-head".into(),
571                serde_json::Value::String("ready".into()),
572            )]),
573        };
574        let body: BodyStream = Box::pin(futures_util::stream::pending());
575        sender
576            .send(ClientDelivery {
577                core: CoreClientDelivery::Terminal(
578                    unb_core::ApplicationFrame::from_envelope(&response).unwrap(),
579                ),
580                body: Some(WireBody::Stream(body)),
581            })
582            .await
583            .unwrap();
584
585        let response =
586            tokio::time::timeout(std::time::Duration::from_millis(50), stream.next_response())
587                .await
588                .expect("the response head must not wait for body completion")
589                .unwrap()
590                .unwrap();
591        assert_eq!(response.head().headers["x-head"], "ready");
592
593        let body = response.into_body();
594        assert!(
595            tokio::time::timeout(std::time::Duration::from_millis(20), body.collect_to(1024),)
596                .await
597                .is_err(),
598            "body consumption remains independently pending"
599        );
600    }
601
602    #[tokio::test]
603    async fn dropping_a_saturated_stream_unblocks_delivery() {
604        let (directives, _directive_receiver) = mpsc::channel(4);
605        let session = ClientSession::connected(directives.clone());
606        let operation = ClientOperationId::from("saturated");
607        let (sender, receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
608        session.register(operation.clone(), sender).unwrap();
609        let stream = ClientStream {
610            operation: operation.clone(),
611            receiver,
612            directives,
613            completed: false,
614            pending: None,
615        };
616
617        let blocked = tokio::spawn({
618            let session = session.clone();
619            let operation = operation.clone();
620            async move {
621                for sequence in 0..CLIENT_OPERATION_QUEUE * 4 {
622                    session
623                        .deliver(
624                            &operation,
625                            CoreClientDelivery::Item(
626                                unb_core::ApplicationFrame::from_envelope(&event(sequence))
627                                    .unwrap(),
628                            ),
629                            None,
630                        )
631                        .await;
632                }
633            }
634        });
635        tokio::time::sleep(std::time::Duration::from_millis(10)).await;
636        assert!(!blocked.is_finished());
637
638        drop(stream);
639
640        tokio::time::timeout(std::time::Duration::from_secs(1), blocked)
641            .await
642            .expect("dropping the consumer must unblock pending delivery")
643            .unwrap();
644        assert!(!session
645            .operations
646            .lock()
647            .expect("client operations lock")
648            .contains_key(&operation));
649    }
650
651    #[tokio::test]
652    async fn a_saturated_directive_queue_still_delivers_the_drop_cancel() {
653        let (directives, mut receiver) = mpsc::channel(1);
654        directives
655            .send(Directive::Control {
656                kind: Kind::Ping,
657                payload: Bytes::new(),
658            })
659            .await
660            .unwrap();
661        let (_deliveries, delivery_receiver) = mpsc::channel(1);
662        let stream = ClientStream {
663            operation: ClientOperationId::from("saturated"),
664            receiver: delivery_receiver,
665            directives,
666            completed: false,
667            pending: None,
668        };
669
670        drop(stream);
671        assert!(matches!(
672            receiver.recv().await,
673            Some(Directive::Control { .. })
674        ));
675        assert!(matches!(
676            tokio::time::timeout(std::time::Duration::from_secs(1), receiver.recv())
677                .await
678                .expect("the deferred cancel must arrive"),
679            Some(Directive::CancelClientOperation { operation }) if operation.as_str() == "saturated"
680        ));
681        assert!(receiver.recv().await.is_none());
682    }
683}