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}