use crate::{
ClientError, ClientResult, Response,
conn::{KeepAlive, Mode, ShortConn},
meta::{BeginRequestRec, EndRequestRec, Header, ParamPairs, RequestType, Role},
params::Params,
request::Request,
response::ResponseStream,
};
use bytes::BytesMut;
use std::{
marker::PhantomData,
ops::{Deref, DerefMut},
};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use tracing::debug;
const REQUEST_ID: u16 = 1;
pub struct Client<S, M> {
stream: S,
_mode: PhantomData<M>,
}
impl<S: AsyncRead + AsyncWrite + Unpin> Client<S, ShortConn> {
pub fn new(stream: S) -> Self {
Self {
stream,
_mode: PhantomData,
}
}
pub async fn execute_once<I: AsyncRead + Unpin>(
mut self,
request: Request<'_, I>,
) -> ClientResult<Response> {
self.inner_execute(request).await
}
pub async fn execute_once_stream<I: AsyncRead + Unpin>(
mut self,
request: Request<'_, I>,
) -> ClientResult<ResponseStream<S>> {
Self::handle_request(&mut self.stream, REQUEST_ID, request.params, request.stdin).await?;
Ok(ResponseStream::new(self.stream, REQUEST_ID))
}
}
impl<S: AsyncRead + AsyncWrite + Unpin> Client<S, KeepAlive> {
pub fn new_keep_alive(stream: S) -> Self {
Self {
stream,
_mode: PhantomData,
}
}
pub async fn execute<I: AsyncRead + Unpin>(
&mut self,
request: Request<'_, I>,
) -> ClientResult<Response> {
self.inner_execute(request).await
}
pub async fn execute_stream<I: AsyncRead + Unpin>(
&mut self,
request: Request<'_, I>,
) -> ClientResult<ResponseStream<&mut S>> {
Self::handle_request(&mut self.stream, REQUEST_ID, request.params, request.stdin).await?;
Ok(ResponseStream::new(&mut self.stream, REQUEST_ID))
}
}
impl<S: AsyncRead + AsyncWrite + Unpin, M: Mode> Client<S, M> {
async fn inner_execute<I: AsyncRead + Unpin>(
&mut self,
request: Request<'_, I>,
) -> ClientResult<Response> {
Self::handle_request(&mut self.stream, REQUEST_ID, request.params, request.stdin).await?;
Self::handle_response(&mut self.stream, REQUEST_ID).await
}
async fn handle_request<'a, I: AsyncRead + Unpin>(
stream: &mut S,
id: u16,
params: Params<'a>,
mut body: I,
) -> ClientResult<()> {
Self::handle_request_start(stream, id).await?;
Self::handle_request_params(stream, id, params).await?;
Self::handle_request_body(stream, id, &mut body).await?;
Self::handle_request_flush(stream).await?;
Ok(())
}
async fn handle_request_start(stream: &mut S, id: u16) -> ClientResult<()> {
debug!(id, "Start handle request");
let begin_request_rec = BeginRequestRec::new(id, Role::Responder, <M>::is_keep_alive());
begin_request_rec.write_to_stream(stream).await?;
Ok(())
}
async fn handle_request_params<'a>(
stream: &mut S,
id: u16,
params: Params<'a>,
) -> ClientResult<()> {
let param_pairs = ParamPairs::new(params);
debug!(id, "Params will be sent {param_pairs:#?}.");
Header::write_to_stream_batches(
RequestType::Params,
id,
stream,
&mut param_pairs.to_content().as_ref(),
Some(|header| {
debug!(id, ?header, "Send to stream for Params.");
header
}),
)
.await?;
Header::write_to_stream_batches(
RequestType::Params,
id,
stream,
&mut tokio::io::empty(),
Some(|header| {
debug!(id, ?header, "Send to stream for Params.");
header
}),
)
.await?;
Ok(())
}
async fn handle_request_body<I: AsyncRead + Unpin>(
stream: &mut S,
id: u16,
body: &mut I,
) -> ClientResult<()> {
Header::write_to_stream_batches(
RequestType::Stdin,
id,
stream,
body,
Some(|header| {
debug!(id, ?header, "Send to stream for Stdin.");
header
}),
)
.await?;
Header::write_to_stream_batches(
RequestType::Stdin,
id,
stream,
&mut tokio::io::empty(),
Some(|header| {
debug!(id, ?header, "Send to stream for Stdin.");
header
}),
)
.await?;
Ok(())
}
async fn handle_request_flush(stream: &mut S) -> ClientResult<()> {
stream.flush().await?;
Ok(())
}
async fn handle_response(stream: &mut S, id: u16) -> ClientResult<Response> {
let mut response = Response::default();
let mut stderr = BytesMut::new();
let mut stdout = BytesMut::new();
loop {
let header = Header::new_from_stream(stream).await?;
if header.request_id != id {
return Err(ClientError::ResponseNotFound { id });
}
debug!(id, ?header, "Receive from stream.");
match header.r#type {
RequestType::Stdout => {
stdout.extend_from_slice(&header.read_content_from_stream(stream).await?);
}
RequestType::Stderr => {
stderr.extend_from_slice(&header.read_content_from_stream(stream).await?);
}
RequestType::EndRequest => {
let end_request_rec = EndRequestRec::from_header(&header, stream).await?;
debug!(id, ?end_request_rec, "Receive from stream.");
end_request_rec
.end_request
.protocol_status
.convert_to_client_result(end_request_rec.end_request.app_status)?;
response.stdout = if stdout.is_empty() {
None
} else {
Some(stdout.freeze())
};
response.stderr = if stderr.is_empty() {
None
} else {
Some(stderr.freeze())
};
return Ok(response);
}
r#type => {
return Err(ClientError::UnknownRequestType {
request_type: r#type,
});
}
}
}
}
}
impl<S, M> Deref for Client<S, M> {
type Target = S;
fn deref(&self) -> &Self::Target {
&self.stream
}
}
impl<S, M> DerefMut for Client<S, M> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.stream
}
}