pub mod error;
use aqueue::Actor;
use error::Result;
use log::*;
use std::borrow::Cow;
use std::future::Future;
use std::net::SocketAddr;
use std::ops::Deref;
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt, ReadHalf, WriteHalf};
use tokio::net::{TcpStream, ToSocketAddrs};
pub struct TcpClient<T> {
disconnect: bool,
sender: WriteHalf<T>,
peer_addr: SocketAddr,
}
impl TcpClient<TcpStream> {
#[inline]
pub async fn connect<
Addr: ToSocketAddrs,
Fut: Future<Output = std::result::Result<bool, E>> + Send + 'static,
E: std::fmt::Display + Send + 'static,
Token: Send + 'static,
>(
addr: Addr,
input: impl FnOnce(Token, Arc<Actor<TcpClient<TcpStream>>>, ReadHalf<TcpStream>) -> Fut
+ Send
+ 'static,
token: Token,
) -> Result<Arc<Actor<TcpClient<TcpStream>>>> {
let stream = TcpStream::connect(addr).await?;
let target = stream.peer_addr()?;
Self::init(input, token, stream, target)
}
#[inline]
pub async fn connect_with_timeout<
Addr: ToSocketAddrs,
Fut: Future<Output = std::result::Result<bool, E>> + Send + 'static,
E: std::fmt::Display + Send + 'static,
Token: Send + 'static,
>(
addr: Addr,
duration: Duration,
input: impl FnOnce(Token, Arc<Actor<TcpClient<TcpStream>>>, ReadHalf<TcpStream>) -> Fut
+ Send
+ 'static,
token: Token,
) -> Result<Arc<Actor<TcpClient<TcpStream>>>> {
let stream = tokio::time::timeout(duration, TcpStream::connect(addr))
.await
.map_err(|_| {
std::io::Error::new(std::io::ErrorKind::TimedOut, "connection timed out")
})??;
let target = stream.peer_addr()?;
Self::init(input, token, stream, target)
}
}
impl<T> TcpClient<T>
where
T: AsyncRead + AsyncWrite + Send + 'static,
{
#[inline]
pub async fn connect_stream_type<
Addr: ToSocketAddrs,
Fut: Future<Output = std::result::Result<bool, E>> + Send + 'static,
E: std::fmt::Display + Send + 'static,
StreamFut: Future<Output = anyhow::Result<T>> + Send + 'static,
Token: Send + 'static,
>(
addr: Addr,
stream_init: impl FnOnce(TcpStream) -> StreamFut + Send + 'static,
input: impl FnOnce(Token, Arc<Actor<TcpClient<T>>>, ReadHalf<T>) -> Fut + Send + 'static,
token: Token,
) -> Result<Arc<Actor<TcpClient<T>>>> {
let stream = TcpStream::connect(addr).await?;
let target = stream.peer_addr()?;
let stream = stream_init(stream).await?;
Self::init(input, token, stream, target)
}
#[inline]
pub async fn connect_stream_type_with_timeout<
Addr: ToSocketAddrs,
Fut: Future<Output = std::result::Result<bool, E>> + Send + 'static,
E: std::fmt::Display + Send + 'static,
StreamFut: Future<Output = anyhow::Result<T>> + Send + 'static,
Token: Send + 'static,
>(
addr: Addr,
duration: Duration,
stream_init: impl FnOnce(TcpStream) -> StreamFut + Send + 'static,
input: impl FnOnce(Token, Arc<Actor<TcpClient<T>>>, ReadHalf<T>) -> Fut + Send + 'static,
token: Token,
) -> Result<Arc<Actor<TcpClient<T>>>> {
let stream = tokio::time::timeout(duration, TcpStream::connect(addr))
.await
.map_err(|_| {
std::io::Error::new(std::io::ErrorKind::TimedOut, "connection timed out")
})??;
let target = stream.peer_addr()?;
let stream = stream_init(stream).await?;
Self::init(input, token, stream, target)
}
#[inline]
fn init<Fut, E, Token>(
f: impl FnOnce(Token, Arc<Actor<TcpClient<T>>>, ReadHalf<T>) -> Fut + Send + 'static,
token: Token,
stream: T,
target: SocketAddr,
) -> Result<Arc<Actor<TcpClient<T>>>>
where
Fut: Future<Output = std::result::Result<bool, E>> + Send + 'static,
E: std::fmt::Display + Send + 'static,
Token: Send + 'static,
{
let (reader, sender) = tokio::io::split(stream);
let client = Arc::new(Actor::new(TcpClient {
disconnect: false,
sender,
peer_addr: target,
}));
let read_client = client.clone();
tokio::spawn(async move {
let disconnect_client = read_client.clone();
let need_disconnect = f(token, read_client, reader).await.unwrap_or_else(|err| {
error!("reader error:{}", err);
true });
if need_disconnect {
if let Err(er) = disconnect_client.disconnect().await {
error!("disconnect to{} err:{}", target, er);
} else {
debug!("disconnect to {}", target);
}
} else {
debug!("{} reader is close", target);
}
});
Ok(client)
}
#[inline]
fn ensure_connected(&self) -> Result<()> {
if self.disconnect {
Err(error::Error::SendError(Cow::Owned(format!(
"Disconnect, peer:{}",
self.peer_addr
))))
} else {
Ok(())
}
}
#[inline]
async fn disconnect(&mut self) -> Result<()> {
if !self.disconnect {
self.sender.shutdown().await?;
self.disconnect = true;
}
Ok(())
}
#[inline]
async fn send(&mut self, buff: &[u8]) -> Result<usize> {
self.ensure_connected()?;
Ok(self.sender.write(buff).await?)
}
#[inline]
async fn send_all(&mut self, buff: &[u8]) -> Result<()> {
self.ensure_connected()?;
self.sender.write_all(buff).await?;
Ok(self.sender.flush().await?)
}
#[inline]
async fn flush(&mut self) -> Result<()> {
self.ensure_connected()?;
Ok(self.sender.flush().await?)
}
}
pub trait SocketClientTrait {
#[must_use]
fn send<B: Deref<Target = [u8]> + Send + 'static>(
&self,
buff: B,
) -> impl Future<Output = Result<usize>>;
fn send_all<B: Deref<Target = [u8]> + Send + 'static>(
&self,
buff: B,
) -> impl Future<Output = Result<()>>;
#[must_use]
fn send_ref(&self, buff: &[u8]) -> impl Future<Output = Result<usize>>;
fn send_all_ref(&self, buff: &[u8]) -> impl Future<Output = Result<()>>;
fn flush(&self) -> impl Future<Output = Result<()>>;
fn disconnect(&self) -> impl Future<Output = Result<()>>;
}
impl<T> SocketClientTrait for Actor<TcpClient<T>>
where
T: AsyncRead + AsyncWrite + Send + 'static,
{
#[inline]
async fn send<B: Deref<Target = [u8]> + Send + 'static>(&self, buff: B) -> Result<usize> {
self.inner_call(|inner| async move { inner.get_mut().send(&buff).await })
.await
}
#[inline]
async fn send_all<B: Deref<Target = [u8]> + Send + 'static>(&self, buff: B) -> Result<()> {
self.inner_call(|inner| async move { inner.get_mut().send_all(&buff).await })
.await
}
#[inline]
async fn send_ref(&self, buff: &[u8]) -> Result<usize> {
if buff.is_empty() {
return Err(error::Error::SendError(Cow::Borrowed("send buff is none")));
}
self.inner_call(|inner| async move { inner.get_mut().send(buff).await })
.await
}
#[inline]
async fn send_all_ref(&self, buff: &[u8]) -> Result<()> {
if buff.is_empty() {
return Err(error::Error::SendError(Cow::Borrowed("send buff is none")));
}
self.inner_call(|inner| async move { inner.get_mut().send_all(buff).await })
.await
}
#[inline]
async fn flush(&self) -> Result<()> {
self.inner_call(|inner| async move { inner.get_mut().flush().await })
.await
}
#[inline]
async fn disconnect(&self) -> Result<()> {
self.inner_call(|inner| async move { inner.get_mut().disconnect().await })
.await
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::net::TcpListener;
#[test]
fn test_send_error_display() {
let err = error::Error::SendError(Cow::Borrowed("test message"));
assert_eq!(err.to_string(), "test message");
}
#[test]
fn test_io_error_conversion() {
let io_err = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused");
let err: error::Error = io_err.into();
assert!(matches!(err, error::Error::IOError(_)));
}
#[tokio::test]
async fn test_echo_client() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (socket, _) = listener.accept().await.unwrap();
let (mut reader, mut writer) = tokio::io::split(socket);
let _ = tokio::io::copy(&mut reader, &mut writer).await;
});
let client = TcpClient::connect(
addr,
async move |_, _client, _reader| {
Ok::<bool, Box<dyn std::error::Error + Send>>(false)
},
(),
)
.await
.unwrap();
client.send_all_ref(b"hello echo").await.unwrap();
client.disconnect().await.unwrap();
}
#[tokio::test]
async fn test_connect_timeout() {
let result = TcpClient::connect_with_timeout(
"192.0.2.1:9999",
Duration::from_millis(100),
async move |_, _client, _reader| Ok::<bool, Box<dyn std::error::Error + Send>>(false),
(),
)
.await;
assert!(result.is_err());
}
}