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}