use core::pin::Pin;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use tokio::net::{TcpStream, ToSocketAddrs};
use crate::alloc::Allocator;
use crate::builder::BlockBuilder;
use crate::client::{ClientOpts, Event, ServerInfo};
use crate::codec::Codec;
use crate::error::{Error, ErrorKind, Result};
use crate::ioless::{IolessClient, Step};
const DEFAULT_READ_BUF_BYTES: usize = 8 * 1024;
pub trait AsyncTransport: AsyncRead + AsyncWrite + Unpin + Send {}
impl<S: AsyncRead + AsyncWrite + Unpin + Send> AsyncTransport for S {}
pub type BoxedAsyncClient = AsyncClient<Box<dyn AsyncTransport>>;
pub struct AsyncClient<S = TcpStream> {
core: IolessClient,
stream: S,
read_buf: Vec<u8>,
}
impl AsyncClient<TcpStream> {
pub async fn connect<A>(
addr: A,
opts: ClientOpts,
codec: Option<Pin<Box<Codec>>>,
) -> Result<Self>
where
A: ToSocketAddrs,
{
let sock = TcpStream::connect(addr).await?;
sock.set_nodelay(true).ok();
Self::handshake_on(sock, opts, codec).await
}
}
#[cfg(feature = "tls")]
impl AsyncClient<tokio_rustls::client::TlsStream<TcpStream>> {
pub async fn connect_tls<A>(
addr: A,
domain: &str,
opts: ClientOpts,
codec: Option<Pin<Box<Codec>>>,
config: std::sync::Arc<rustls::ClientConfig>,
) -> Result<Self>
where
A: ToSocketAddrs,
{
let sock = TcpStream::connect(addr).await?;
sock.set_nodelay(true).ok();
let server_name =
rustls::pki_types::ServerName::try_from(domain.to_owned()).map_err(|_| {
Error::new(
ErrorKind::Usage,
format!("invalid TLS server name: {domain}"),
)
})?;
let tls = tokio_rustls::TlsConnector::from(config)
.connect(server_name, sock)
.await
.map_err(|e| Error::new(ErrorKind::Io, format!("TLS handshake: {e}")))?;
Self::handshake_on(tls, opts, codec).await
}
}
impl<S: AsyncTransport> AsyncClient<S> {
pub async fn handshake_on(
stream: S,
opts: ClientOpts,
codec: Option<Pin<Box<Codec>>>,
) -> Result<Self> {
let read_buf_bytes = if opts.read_buffer_bytes == 0 {
DEFAULT_READ_BUF_BYTES
} else {
opts.read_buffer_bytes
};
let mut client = Self {
core: IolessClient::new(&opts, Allocator::stdlib(), codec)?,
stream,
read_buf: vec![0; read_buf_bytes],
};
client.pump_until_ready(|core| core.handshake()).await?;
Ok(client)
}
pub async fn send_query(&mut self, sql: &str, query_id: Option<&str>) -> Result<()> {
self.drain_out().await?;
self.core.send_query(sql, query_id)?;
self.drain_out().await
}
pub async fn send_data(&mut self, builder: Option<&BlockBuilder<'_>>) -> Result<()> {
self.drain_out().await?;
self.core.send_data(builder)?;
self.drain_out().await
}
pub async fn send_data_end(&mut self) -> Result<()> {
self.drain_out().await?;
self.core.send_data_end()?;
self.drain_out().await
}
pub async fn recv_event(&mut self) -> Result<Event> {
let mut event = None;
self.pump_until_ready(|core| {
Ok(match core.recv_event()? {
Step::Ready(e) => {
event = Some(e);
Step::Ready(())
}
Step::NeedsInput => Step::NeedsInput,
})
})
.await?;
Ok(event.expect("pump_until_ready only returns once the step stored an event"))
}
pub fn server_info(&self) -> Option<ServerInfo> {
self.core.server_info()
}
pub fn boxed(self) -> BoxedAsyncClient
where
S: 'static,
{
AsyncClient {
core: self.core,
stream: Box::new(self.stream),
read_buf: self.read_buf,
}
}
pub fn core(&mut self) -> &mut IolessClient {
&mut self.core
}
async fn drain_out(&mut self) -> Result<()> {
let mut wrote = false;
loop {
let buf = self.core.pending_out();
if buf.is_empty() {
break;
}
let n = self.stream.write(buf).await?;
if n == 0 {
return Err(Error::new(ErrorKind::Io, "transport write returned zero"));
}
self.core.consume_out(n);
wrote = true;
}
if wrote {
self.stream.flush().await?;
}
Ok(())
}
async fn pump_until_ready(
&mut self,
mut step: impl FnMut(&mut IolessClient) -> Result<Step<()>>,
) -> Result<()> {
loop {
self.drain_out().await?;
match step(&mut self.core)? {
Step::Ready(()) => return self.drain_out().await,
Step::NeedsInput => {
self.drain_out().await?;
self.read_more().await?;
}
}
}
}
async fn read_more(&mut self) -> Result<()> {
let n = self.stream.read(&mut self.read_buf).await?;
if n == 0 {
return Err(Error::new(ErrorKind::Eof, "transport closed"));
}
self.core.submit(&self.read_buf[..n])
}
}
#[cfg(test)]
mod tests {
use super::{AsyncClient, Event};
#[test]
fn async_client_is_send() {
fn assert_send<T: Send>() {}
assert_send::<AsyncClient>();
assert_send::<Event>();
}
}