use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use bytes::Bytes;
use serde_json::Value;
use tokio::sync::mpsc;
use unb_core::{
ClientDelivery as CoreClientDelivery, ClientOperationId, DiscoverPlan, Envelope, ErrorCode,
Kind, TargetPath,
};
use crate::wire::Directive;
use crate::{BodyStream, WireBody};
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum ClientError {
#[error("client operation {0} is gone")]
Gone(String),
#[error("client operation {operation:?} failed with {code:?}: {message}")]
Protocol {
operation: ClientOperationId,
code: ErrorCode,
message: String,
},
#[error("client operation {0} was cancelled")]
Cancelled(String),
#[error("client operation {0} timed out")]
Timeout(String),
#[error("client session closed during operation {0}")]
SessionClosed(String),
#[error("invalid client request: {0}")]
Invalid(String),
}
const CLIENT_OPERATION_QUEUE: usize = 64;
#[derive(Clone)]
pub struct ClientSession {
operations: Arc<Mutex<HashMap<ClientOperationId, mpsc::Sender<ClientDelivery>>>>,
directives: mpsc::Sender<Directive>,
}
impl ClientSession {
pub(crate) fn connected(directives: mpsc::Sender<Directive>) -> Self {
Self {
operations: Arc::default(),
directives,
}
}
pub async fn start(
&self,
target_path: &str,
kind: Kind,
payload: Bytes,
hops: Option<u8>,
headers: serde_json::Map<String, Value>,
) -> Result<ClientStream, ClientError> {
self.start_with_timeout(target_path, kind, payload, hops, headers, None, None)
.await
}
pub async fn start_streaming(
&self,
target_path: &str,
kind: Kind,
body: BodyStream,
hops: Option<u8>,
headers: serde_json::Map<String, Value>,
) -> Result<ClientStream, ClientError> {
self.start_with_timeout(
target_path,
kind,
Bytes::new(),
hops,
headers,
Some(body),
None,
)
.await
}
pub async fn fetch(
&self,
request: http::Request<Bytes>,
timeout: std::time::Duration,
) -> Result<http::Response<Bytes>, ClientError> {
let (parts, body) = request.into_parts();
let response = self
.fetch_body(
http::Request::from_parts(parts, WireBody::Bytes(body)),
timeout,
)
.await?;
let (parts, body) = response.into_parts();
let body = body
.collect_to(unb_transport::DEFAULT_MAX_FRAME_SIZE)
.await
.map_err(|error| ClientError::Invalid(error.to_string()))?;
Ok(http::Response::from_parts(parts, body))
}
pub async fn fetch_body(
&self,
request: http::Request<WireBody>,
timeout: std::time::Duration,
) -> Result<http::Response<WireBody>, ClientError> {
let (parts, body) = request.into_parts();
let envelope = Envelope::from_request(http::Request::from_parts(parts, Bytes::new()))
.map_err(|error| ClientError::Invalid(error.to_string()))?;
let (payload, body) = match body {
WireBody::Bytes(payload) => (payload, None),
WireBody::Stream(body) => (Bytes::new(), Some(body)),
};
let mut stream = self
.start_with_timeout(
&TargetPath::application(&envelope.target, &envelope.subject)
.map_err(|error| ClientError::Invalid(error.to_string()))?
.to_string(),
Kind::Request,
payload,
envelope.hops,
envelope.headers,
body,
Some(timeout),
)
.await?;
let response = stream
.next_response()
.await?
.ok_or_else(|| ClientError::Gone(stream.operation().as_str().to_owned()))?;
response
.into_response()
.map_err(|error| ClientError::Invalid(error.to_string()))
}
pub async fn subscribe(
&self,
request: http::Request<Bytes>,
timeout: Option<std::time::Duration>,
) -> Result<ClientStream, ClientError> {
let envelope = Envelope::from_request(request)
.map_err(|error| ClientError::Invalid(error.to_string()))?;
self.start_with_timeout(
&TargetPath::application(&envelope.target, &envelope.subject)
.map_err(|error| ClientError::Invalid(error.to_string()))?
.to_string(),
Kind::Subscribe,
envelope.payload,
envelope.hops,
envelope.headers,
None,
timeout,
)
.await
}
pub async fn discover(
&self,
target_path: &str,
plan: DiscoverPlan,
) -> Result<ClientStream, ClientError> {
let timeout = plan.timeout_ms.map(std::time::Duration::from_millis);
let hops = Some(plan.hops);
let payload = Envelope::encode_payload(
&serde_json::to_value(plan).map_err(|error| ClientError::Invalid(error.to_string()))?,
);
self.start_with_timeout(
target_path,
Kind::Discover,
payload,
hops,
Default::default(),
None,
timeout,
)
.await
}
#[allow(clippy::too_many_arguments)]
async fn start_with_timeout(
&self,
target_path: &str,
kind: Kind,
payload: Bytes,
hops: Option<u8>,
headers: serde_json::Map<String, Value>,
body: Option<BodyStream>,
timeout: Option<std::time::Duration>,
) -> Result<ClientStream, ClientError> {
let (sender, receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
let (reply, response) = tokio::sync::oneshot::channel();
self.directives
.send(Directive::StartClientOperation {
target_path: target_path.to_owned(),
kind,
payload,
hops,
headers,
body,
timeout,
sender,
reply,
})
.await
.map_err(|_| ClientError::Gone("session".into()))?;
let operation = response
.await
.map_err(|_| ClientError::Gone("session".into()))?
.map_err(|error| ClientError::Invalid(error.to_string()))?;
Ok(ClientStream {
operation,
receiver,
directives: self.directives.clone(),
completed: false,
pending: None,
})
}
pub(crate) fn register(
&self,
operation: ClientOperationId,
sender: mpsc::Sender<ClientDelivery>,
) -> Result<(), ClientError> {
let mut operations = self.operations.lock().expect("client operations lock");
if operations.insert(operation.clone(), sender).is_some() {
return Err(ClientError::Gone(operation.as_str().to_owned()));
}
Ok(())
}
pub(crate) async fn deliver(
&self,
operation: &ClientOperationId,
delivery: CoreClientDelivery,
body: Option<WireBody>,
) {
let terminal = matches!(
delivery,
CoreClientDelivery::Terminal(_)
| CoreClientDelivery::Cancelled
| CoreClientDelivery::TimedOut
| CoreClientDelivery::SessionClosed
);
let sender = self
.operations
.lock()
.expect("client operations lock")
.get(operation)
.cloned();
let Some(sender) = sender else { return };
let delivered = sender
.send(ClientDelivery {
core: delivery,
body,
})
.await
.is_ok();
if terminal || !delivered {
self.abandon(operation);
}
}
pub(crate) fn abandon(&self, operation: &ClientOperationId) {
self.operations
.lock()
.expect("client operations lock")
.remove(operation);
}
}
#[doc(hidden)]
pub struct ClientDelivery {
core: CoreClientDelivery,
body: Option<WireBody>,
}
pub struct ClientStream {
operation: ClientOperationId,
receiver: mpsc::Receiver<ClientDelivery>,
directives: mpsc::Sender<Directive>,
completed: bool,
pending: Option<ClientDelivery>,
}
impl ClientStream {
pub fn operation(&self) -> &ClientOperationId {
&self.operation
}
pub async fn next(&mut self) -> Result<Option<Envelope>, ClientError> {
if self.completed {
return Ok(None);
}
let delivery = match self.pending.take() {
Some(delivery) => Some(delivery),
None => self.receiver.recv().await,
};
match delivery {
Some(delivery) => {
let response = self.project_head(delivery)?;
match response {
Some(response) => {
let (mut envelope, body) = response.into_parts();
envelope.payload = body
.collect_to(unb_transport::DEFAULT_MAX_FRAME_SIZE)
.await
.map_err(|error| ClientError::Invalid(error.to_string()))?;
Ok(Some(envelope))
}
None => Ok(None),
}
}
None => {
self.completed = true;
Err(ClientError::Gone(self.operation.as_str().to_owned()))
}
}
}
pub fn try_next(&mut self) -> Result<Option<Option<Envelope>>, ClientError> {
if self.completed {
return Ok(Some(None));
}
match self.receiver.try_recv() {
Ok(mut delivery) => match delivery.body.take() {
Some(WireBody::Stream(stream)) => {
delivery.body = Some(WireBody::Stream(stream));
self.pending = Some(delivery);
Ok(None)
}
Some(WireBody::Bytes(payload)) => {
let mut envelope = self.project_without_body(delivery)?;
if let Some(envelope) = &mut envelope {
envelope.payload = payload;
}
Ok(Some(envelope))
}
None => self.project_without_body(delivery).map(Some),
},
Err(mpsc::error::TryRecvError::Empty) => Ok(None),
Err(mpsc::error::TryRecvError::Disconnected) => {
self.completed = true;
Err(ClientError::Gone(self.operation.as_str().to_owned()))
}
}
}
pub async fn next_response(&mut self) -> Result<Option<ClientResponse>, ClientError> {
if self.completed {
return Ok(None);
}
let delivery = match self.pending.take() {
Some(delivery) => Some(delivery),
None => self.receiver.recv().await,
};
match delivery {
Some(delivery) => self.project_head(delivery),
None => {
self.completed = true;
Err(ClientError::Gone(self.operation.as_str().to_owned()))
}
}
}
fn project_without_body(
&mut self,
delivery: ClientDelivery,
) -> Result<Option<Envelope>, ClientError> {
self.project_head(delivery)
.map(|response| response.map(|response| response.envelope))
}
fn project_head(
&mut self,
delivery: ClientDelivery,
) -> Result<Option<ClientResponse>, ClientError> {
match delivery.core {
CoreClientDelivery::Terminal(frame) if frame.head.kind == Kind::Error => {
self.completed = true;
let error = frame.head.error.unwrap_or(unb_core::ApplicationError {
code: ErrorCode::Protocol,
message: "protocol error".to_owned(),
});
Err(ClientError::Protocol {
operation: self.operation.clone(),
code: error.code,
message: error.message,
})
}
CoreClientDelivery::Terminal(frame) => {
self.completed = true;
Ok(Some(ClientResponse {
envelope: frame.into_envelope(),
body: delivery
.body
.unwrap_or_else(|| WireBody::Bytes(Bytes::new())),
}))
}
CoreClientDelivery::Item(frame) => Ok(Some(ClientResponse {
envelope: frame.into_envelope(),
body: delivery
.body
.unwrap_or_else(|| WireBody::Bytes(Bytes::new())),
})),
CoreClientDelivery::Cancelled => {
self.completed = true;
Err(ClientError::Cancelled(self.operation.as_str().to_owned()))
}
CoreClientDelivery::TimedOut => {
self.completed = true;
Err(ClientError::Timeout(self.operation.as_str().to_owned()))
}
CoreClientDelivery::SessionClosed => {
self.completed = true;
Err(ClientError::SessionClosed(
self.operation.as_str().to_owned(),
))
}
}
}
}
pub struct ClientResponse {
envelope: Envelope,
body: WireBody,
}
impl ClientResponse {
pub fn head(&self) -> &Envelope {
&self.envelope
}
pub fn into_parts(self) -> (Envelope, WireBody) {
(self.envelope, self.body)
}
pub fn into_body(self) -> WireBody {
self.body
}
pub fn into_response(self) -> Result<http::Response<WireBody>, unb_core::CoreError> {
let response = self.envelope.to_response()?;
let (parts, _) = response.into_parts();
Ok(http::Response::from_parts(parts, self.body))
}
}
impl Drop for ClientStream {
fn drop(&mut self) {
if self.completed {
return;
}
self.completed = true;
let command = Directive::CancelClientOperation {
operation: self.operation.clone(),
};
if let Err(mpsc::error::TrySendError::Full(command)) = self.directives.try_send(command) {
let directives = self.directives.clone();
n0_future::task::spawn(async move {
let _ = directives.send(command).await;
});
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn event(sequence: usize) -> Envelope {
Envelope {
v: 1,
id: sequence.to_string(),
target: String::new(),
subject: String::new(),
kind: Kind::Event,
corr: None,
seq: None,
hops: None,
body_token: None,
payload: Bytes::new(),
path: Vec::new(),
headers: Default::default(),
}
}
#[tokio::test]
async fn saturated_delivery_blocks_without_loss_until_capacity_returns() {
let (directives, _directive_receiver) = mpsc::channel(1);
let session = ClientSession::connected(directives);
let operation = ClientOperationId::from("full");
let (sender, mut receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
session.register(operation.clone(), sender).unwrap();
let total = CLIENT_OPERATION_QUEUE * 4;
let producer = tokio::spawn({
let session = session.clone();
let operation = operation.clone();
async move {
for sequence in 0..total {
session
.deliver(
&operation,
CoreClientDelivery::Item(
unb_core::ApplicationFrame::from_envelope(&event(sequence))
.unwrap(),
),
None,
)
.await;
}
let terminal = Envelope {
v: 1,
id: "terminal".into(),
target: String::new(),
subject: String::new(),
kind: Kind::Error,
corr: Some(operation.as_str().to_owned()),
seq: None,
hops: None,
body_token: None,
payload: Envelope::encode_payload(&serde_json::json!({
"code": ErrorCode::Busy,
"message": "busy"
})),
path: Vec::new(),
headers: Default::default(),
};
session
.deliver(
&operation,
CoreClientDelivery::Terminal(
unb_core::ApplicationFrame::from_envelope(&terminal).unwrap(),
),
None,
)
.await;
}
});
for sequence in 0..total {
match receiver.recv().await.unwrap().core {
CoreClientDelivery::Item(frame) => {
assert_eq!(frame.head.id, sequence.to_string())
}
_ => panic!("expected ordered item {sequence}"),
}
}
assert!(matches!(
receiver.recv().await,
Some(ClientDelivery {
core: CoreClientDelivery::Terminal(frame),
..
}) if frame.head.error.as_ref().is_some_and(|error| error.code == ErrorCode::Busy)
));
producer.await.unwrap();
assert!(!session
.operations
.lock()
.expect("client operations lock")
.contains_key(&operation));
}
#[tokio::test]
async fn response_head_is_delivered_before_its_body_completes() {
let (directives, _directive_receiver) = mpsc::channel(1);
let (sender, receiver) = mpsc::channel(1);
let mut stream = ClientStream {
operation: ClientOperationId::from("head-first"),
receiver,
directives,
completed: false,
pending: None,
};
let response = Envelope {
v: 1,
id: "response".into(),
target: String::new(),
subject: String::new(),
kind: Kind::Response,
corr: Some("head-first".into()),
seq: None,
hops: None,
body_token: Some("body-1".into()),
payload: Bytes::new(),
path: Vec::new(),
headers: serde_json::Map::from_iter([(
"x-head".into(),
serde_json::Value::String("ready".into()),
)]),
};
let body: BodyStream = Box::pin(futures_util::stream::pending());
sender
.send(ClientDelivery {
core: CoreClientDelivery::Terminal(
unb_core::ApplicationFrame::from_envelope(&response).unwrap(),
),
body: Some(WireBody::Stream(body)),
})
.await
.unwrap();
let response =
tokio::time::timeout(std::time::Duration::from_millis(50), stream.next_response())
.await
.expect("the response head must not wait for body completion")
.unwrap()
.unwrap();
assert_eq!(response.head().headers["x-head"], "ready");
let body = response.into_body();
assert!(
tokio::time::timeout(std::time::Duration::from_millis(20), body.collect_to(1024),)
.await
.is_err(),
"body consumption remains independently pending"
);
}
#[tokio::test]
async fn dropping_a_saturated_stream_unblocks_delivery() {
let (directives, _directive_receiver) = mpsc::channel(4);
let session = ClientSession::connected(directives.clone());
let operation = ClientOperationId::from("saturated");
let (sender, receiver) = mpsc::channel(CLIENT_OPERATION_QUEUE);
session.register(operation.clone(), sender).unwrap();
let stream = ClientStream {
operation: operation.clone(),
receiver,
directives,
completed: false,
pending: None,
};
let blocked = tokio::spawn({
let session = session.clone();
let operation = operation.clone();
async move {
for sequence in 0..CLIENT_OPERATION_QUEUE * 4 {
session
.deliver(
&operation,
CoreClientDelivery::Item(
unb_core::ApplicationFrame::from_envelope(&event(sequence))
.unwrap(),
),
None,
)
.await;
}
}
});
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
assert!(!blocked.is_finished());
drop(stream);
tokio::time::timeout(std::time::Duration::from_secs(1), blocked)
.await
.expect("dropping the consumer must unblock pending delivery")
.unwrap();
assert!(!session
.operations
.lock()
.expect("client operations lock")
.contains_key(&operation));
}
#[tokio::test]
async fn a_saturated_directive_queue_still_delivers_the_drop_cancel() {
let (directives, mut receiver) = mpsc::channel(1);
directives
.send(Directive::Control {
kind: Kind::Ping,
payload: Bytes::new(),
})
.await
.unwrap();
let (_deliveries, delivery_receiver) = mpsc::channel(1);
let stream = ClientStream {
operation: ClientOperationId::from("saturated"),
receiver: delivery_receiver,
directives,
completed: false,
pending: None,
};
drop(stream);
assert!(matches!(
receiver.recv().await,
Some(Directive::Control { .. })
));
assert!(matches!(
tokio::time::timeout(std::time::Duration::from_secs(1), receiver.recv())
.await
.expect("the deferred cancel must arrive"),
Some(Directive::CancelClientOperation { operation }) if operation.as_str() == "saturated"
));
assert!(receiver.recv().await.is_none());
}
}