use crate::MontycatClientError;
use crate::engine::structure::Engine;
#[cfg(feature = "tls")]
use rustls_pki_types::ServerName;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader, ReadBuf};
use tokio::net::TcpStream;
use tokio::sync::watch::Receiver;
use tokio::time::timeout;
#[cfg(feature = "tls")]
use tokio_rustls::TlsConnector;
#[cfg(feature = "tls")]
use tokio_rustls::{
client::TlsStream,
rustls::{ClientConfig, RootCertStore},
};
pub(crate) type StreamCallback = Arc<dyn Fn(&mut [u8]) + Send + Sync>;
const CHUNK_SIZE: usize = 1024 * 256;
pub(crate) enum Connection {
Plain(TcpStream),
#[cfg(feature = "tls")]
Tls(Box<TlsStream<TcpStream>>),
}
impl Connection {
pub(crate) fn split(
self,
) -> (
Box<dyn AsyncRead + Unpin + Send>,
Box<dyn AsyncWrite + Unpin + Send>,
) {
match self {
Connection::Plain(stream) => {
let (r, w) = tokio::io::split(stream);
(Box::new(r), Box::new(w))
}
#[cfg(feature = "tls")]
Connection::Tls(stream) => {
let (r, w) = tokio::io::split(stream);
(Box::new(r), Box::new(w))
}
}
}
pub(crate) async fn shutdown(&mut self) -> Result<(), MontycatClientError> {
match self {
Connection::Plain(stream) => stream.shutdown().await,
#[cfg(feature = "tls")]
Connection::Tls(stream) => stream.shutdown().await,
}
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))
}
}
impl AsyncRead for Connection {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Connection::Plain(stream) => Pin::new(stream).poll_read(cx, buf),
#[cfg(feature = "tls")]
Connection::Tls(stream) => Pin::new(stream.as_mut()).poll_read(cx, buf),
}
}
}
impl AsyncWrite for Connection {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
match self.get_mut() {
Connection::Plain(stream) => Pin::new(stream).poll_write(cx, buf),
#[cfg(feature = "tls")]
Connection::Tls(stream) => Pin::new(stream.as_mut()).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Connection::Plain(stream) => Pin::new(stream).poll_flush(cx),
#[cfg(feature = "tls")]
Connection::Tls(stream) => Pin::new(stream.as_mut()).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Connection::Plain(stream) => Pin::new(stream).poll_shutdown(cx),
#[cfg(feature = "tls")]
Connection::Tls(stream) => Pin::new(stream.as_mut()).poll_shutdown(cx),
}
}
}
pub(crate) async fn send_data(
engine: &Engine,
query: &[u8],
callback: Option<StreamCallback>,
stop_event: Option<&mut Receiver<bool>>,
port_override: Option<u16>,
) -> Result<Option<Vec<u8>>, MontycatClientError> {
let port: u16 = port_override.unwrap_or(engine.port);
match callback {
Some(cb) => subscription(engine, port, query, cb, stop_event).await,
None => request(engine, port, query).await,
}
}
async fn connect(engine: &Engine, port: u16) -> Result<Connection, MontycatClientError> {
let use_tls: bool = engine.use_tls;
let host: String = engine.host.clone();
let plain_stream: TcpStream = TcpStream::connect((host.as_ref(), port))
.await
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
let connection = if use_tls {
#[cfg(feature = "tls")]
{
let mut root_cert_store = RootCertStore::empty();
root_cert_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
root_cert_store.add_parsable_certificates(engine.tls_root_certificates.iter().cloned());
let config = ClientConfig::builder()
.with_root_certificates(root_cert_store)
.with_no_client_auth();
let connector = TlsConnector::from(Arc::new(config));
let server_name = ServerName::try_from(host)
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
let tls_stream = match timeout(
Duration::from_secs(10),
connector.connect(server_name, plain_stream),
)
.await
{
Ok(Ok(stream)) => stream,
Ok(Err(e)) => {
return Err(MontycatClientError::ClientEngineError(format!(
"TLS handshake failed: {}",
e
)));
}
Err(_) => {
return Err(MontycatClientError::ClientEngineError(
"TLS handshake timed out".to_string(),
));
}
};
Connection::Tls(Box::new(tls_stream))
}
#[cfg(not(feature = "tls"))]
{
return Err(MontycatClientError::ClientEngineError(
"TLS feature not enabled".to_string(),
));
}
} else {
Connection::Plain(plain_stream)
};
Ok(connection)
}
async fn subscription(
engine: &Engine,
port: u16,
query: &[u8],
callback: StreamCallback,
stop_event: Option<&mut Receiver<bool>>,
) -> Result<Option<Vec<u8>>, MontycatClientError> {
let (reader, mut writer) = connect(engine, port).await?.split();
let mut reader = BufReader::with_capacity(CHUNK_SIZE, reader);
writer
.write_all(query)
.await
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
writer
.flush()
.await
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
let mut buf = vec![];
loop {
if let Some(ref stop) = stop_event
&& let Ok(true) = stop.has_changed()
&& *stop.borrow()
{
break;
}
buf.clear();
let n = reader
.read_until(b'\n', &mut buf)
.await
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
if n == 0 {
break;
}
callback(buf.as_mut_slice());
}
writer
.shutdown()
.await
.map_err(|e| MontycatClientError::ClientEngineError(e.to_string()))?;
Ok(None)
}
async fn request(
engine: &Engine,
port: u16,
query: &[u8],
) -> Result<Option<Vec<u8>>, MontycatClientError> {
let pool = engine.pool.clone();
if let Some(pool) = &pool
&& let Some(reader) = pool.checkout().await
{
match exchange(reader, query).await {
Ok((reader, response)) => {
pool.checkin(reader).await;
return Ok(Some(response));
}
Err(Exchange::Write(_)) => {
}
Err(Exchange::Read(e)) => {
return Err(e);
}
}
}
let reader = BufReader::with_capacity(CHUNK_SIZE, connect(engine, port).await?);
match exchange(reader, query).await {
Ok((mut reader, response)) => {
match &pool {
Some(pool) => pool.checkin(reader).await,
None => {
reader.get_mut().shutdown().await?;
}
}
Ok(Some(response))
}
Err(Exchange::Write(e) | Exchange::Read(e)) => Err(e),
}
}
enum Exchange {
Write(MontycatClientError),
Read(MontycatClientError),
}
async fn exchange(
mut reader: BufReader<Connection>,
query: &[u8],
) -> Result<(BufReader<Connection>, Vec<u8>), Exchange> {
let write = async {
reader.get_mut().write_all(query).await?;
reader.get_mut().flush().await
}
.await;
if let Err(e) = write {
return Err(Exchange::Write(MontycatClientError::ClientEngineError(
e.to_string(),
)));
}
let mut buf = vec![];
let read = timeout(Duration::from_secs(120), reader.read_until(b'\n', &mut buf)).await;
match read {
Err(e) => Err(Exchange::Read(MontycatClientError::ClientEngineError(
e.to_string(),
))),
Ok(Err(e)) => Err(Exchange::Read(MontycatClientError::ClientEngineError(
e.to_string(),
))),
Ok(Ok(0)) => Err(Exchange::Read(MontycatClientError::ClientEngineError(
"connection closed before a response was received".to_string(),
))),
Ok(Ok(_)) => Ok((reader, buf)),
}
}