Skip to main content

unb_runtime/
client.rs

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