use super::stream::GlommioStream;
use crate::net::yolo_accept;
use crate::parking::Reactor;
use crate::GlommioError;
use crate::Local;
use futures_lite::future::poll_fn;
use futures_lite::io::{AsyncBufRead, AsyncRead, AsyncWrite};
use futures_lite::ready;
use futures_lite::stream::{self, Stream};
use nix::sys::socket::{InetAddr, SockAddr};
use pin_project_lite::pin_project;
use socket2::{Domain, Protocol, Socket, Type};
use std::io;
use std::net::{self, Shutdown, SocketAddr, ToSocketAddrs};
use std::os::unix::io::RawFd;
use std::os::unix::io::{AsRawFd, FromRawFd};
use std::pin::Pin;
use std::rc::{Rc, Weak};
use std::task::{Context, Poll};
type Result<T> = crate::Result<T, ()>;
#[derive(Debug)]
pub struct TcpListener {
reactor: Weak<Reactor>,
listener: net::TcpListener,
}
impl TcpListener {
pub fn bind<A: ToSocketAddrs>(addr: A) -> Result<TcpListener> {
let addr = addr
.to_socket_addrs()
.unwrap()
.next()
.ok_or_else(|| io::Error::new(io::ErrorKind::Other, "empty address"))?;
let domain = if addr.is_ipv6() {
Domain::ipv6()
} else {
Domain::ipv4()
};
let sk = Socket::new(domain, Type::stream(), Some(Protocol::tcp()))?;
let addr = socket2::SockAddr::from(addr);
sk.set_reuse_port(true)?;
sk.bind(&addr)?;
sk.listen(1024)?;
let listener = sk.into_tcp_listener();
Ok(TcpListener {
reactor: Rc::downgrade(&Local::get_reactor()),
listener,
})
}
pub async fn shared_accept(&self) -> Result<AcceptedTcpStream> {
let reactor = self.reactor.upgrade().unwrap();
let raw_fd = self.listener.as_raw_fd();
if let Some(r) = yolo_accept(raw_fd) {
match r {
Ok(fd) => {
return Ok(AcceptedTcpStream { fd });
}
Err(err) => return Err(GlommioError::IoError(err)),
}
}
let source = reactor.accept(self.listener.as_raw_fd());
let fd = source.collect_rw().await?;
Ok(AcceptedTcpStream { fd: fd as RawFd })
}
pub async fn accept(&self) -> Result<TcpStream> {
let a = self.shared_accept().await?;
Ok(a.bind_to_executor())
}
pub fn incoming(&self) -> impl Stream<Item = Result<TcpStream>> + Unpin + '_ {
Box::pin(stream::unfold(self, |listener| async move {
let res = listener.accept().await;
Some((res, listener))
}))
}
pub fn local_addr(&self) -> Result<SocketAddr> {
Ok(self.listener.local_addr()?)
}
pub fn ttl(&self) -> Result<u32> {
Ok(self.listener.ttl()?)
}
pub fn set_ttl(&self, ttl: u32) -> Result<()> {
Ok(self.listener.set_ttl(ttl)?)
}
}
#[derive(Copy, Clone, Debug)]
pub struct AcceptedTcpStream {
fd: RawFd,
}
impl AcceptedTcpStream {
pub fn bind_to_executor(self) -> TcpStream {
TcpStream {
stream: unsafe { GlommioStream::from_raw_fd(self.fd) },
}
}
}
pin_project! {
#[derive(Debug)]
pub struct TcpStream {
stream: GlommioStream<net::TcpStream>
}
}
impl From<socket2::Socket> for TcpStream {
fn from(socket: socket2::Socket) -> TcpStream {
Self {
stream: GlommioStream::<net::TcpStream>::from(socket),
}
}
}
impl AsRawFd for TcpStream {
fn as_raw_fd(&self) -> RawFd {
self.stream.as_raw_fd()
}
}
impl FromRawFd for TcpStream {
unsafe fn from_raw_fd(fd: RawFd) -> Self {
let socket = socket2::Socket::from_raw_fd(fd);
TcpStream::from(socket)
}
}
impl TcpStream {
pub async fn connect<A: ToSocketAddrs>(addr: A) -> Result<TcpStream> {
let addr = addr.to_socket_addrs()?.next().unwrap();
let reactor = Local::get_reactor();
let domain = if addr.is_ipv6() {
Domain::ipv6()
} else {
Domain::ipv4()
};
let socket = Socket::new(domain, Type::stream(), Some(Protocol::tcp()))?;
let inet = InetAddr::from_std(&addr);
let addr = SockAddr::new_inet(inet);
let source = reactor.connect(socket.as_raw_fd(), addr);
source.collect_rw().await?;
Ok(TcpStream {
stream: GlommioStream::from(socket),
})
}
pub async fn shutdown(&self, how: Shutdown) -> Result<()> {
poll_fn(|cx| self.stream.poll_shutdown(cx, how))
.await
.map_err(Into::into)
}
pub fn set_nodelay(&self, value: bool) -> Result<()> {
self.stream.stream.set_nodelay(value).map_err(Into::into)
}
pub fn nodelay(&self) -> Result<bool> {
self.stream.stream.nodelay().map_err(Into::into)
}
pub fn set_buffer_size(&mut self, buffer_size: usize) {
self.stream.rx_buf_size = buffer_size;
}
pub fn buffer_size(&mut self) -> usize {
self.stream.rx_buf_size
}
pub fn ttl(&self) -> Result<u32> {
Ok(self.stream.stream.ttl()?)
}
pub fn set_ttl(&self, ttl: u32) -> Result<()> {
Ok(self.stream.stream.set_ttl(ttl)?)
}
pub async fn peek(&self, buf: &mut [u8]) -> Result<usize> {
self.stream.peek(buf).await.map_err(Into::into)
}
pub fn peer_addr(&self) -> Result<SocketAddr> {
self.stream.stream.peer_addr().map_err(Into::into)
}
pub fn local_addr(&self) -> Result<SocketAddr> {
self.stream.stream.local_addr().map_err(Into::into)
}
}
impl AsyncBufRead for TcpStream {
fn poll_fill_buf<'a>(
mut self: Pin<&'a mut Self>,
cx: &mut Context<'_>,
) -> Poll<io::Result<&'a [u8]>> {
let buf_size = self.stream.rx_buf_size;
if self.stream.rx_buf.as_ref().is_none() {
poll_err!(ready!(self.stream.poll_replenish_buffer(cx, buf_size)));
}
let this = self.project();
Poll::Ready(Ok(this.stream.rx_buf.as_ref().unwrap().as_bytes()))
}
fn consume(mut self: Pin<&mut Self>, amt: usize) {
Pin::new(&mut self.stream).consume(amt)
}
}
impl AsyncRead for TcpStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.stream).poll_read(cx, buf)
}
}
impl AsyncWrite for TcpStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.stream).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.stream).poll_flush(cx)
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.stream).poll_close(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::channels::shared_channel;
use crate::enclose;
use crate::timer::Timer;
use crate::LocalExecutorBuilder;
use futures_lite::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt};
use futures_lite::StreamExt;
use std::cell::Cell;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
#[test]
fn tcp_listener_ttl() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_ttl(100).unwrap();
assert_eq!(listener.ttl().unwrap(), 100);
});
}
#[test]
fn tcp_stream_ttl() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let stream = TcpStream::connect(addr).await.unwrap();
stream.set_ttl(100).unwrap();
assert_eq!(stream.ttl().unwrap(), 100);
});
}
#[test]
fn tcp_stream_nodelay() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let stream = TcpStream::connect(addr).await.unwrap();
stream.set_nodelay(true).expect("set_nodelay call failed");
assert_eq!(stream.nodelay().unwrap(), true);
});
}
#[test]
fn connect_local_server() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let coord = Rc::new(Cell::new(0));
let listener_handle: Task<Result<SocketAddr>> =
Task::local(enclose! { (coord) async move {
coord.set(1);
let stream = listener.accept().await?;
Ok(stream.peer_addr()?)
}});
while coord.get() != 1 {
Local::later().await;
}
let stream = TcpStream::connect(addr).await.unwrap();
assert_eq!(listener_handle.await.unwrap(), stream.local_addr().unwrap());
});
}
#[test]
fn multi_executor_bind_works() {
test_executor!(async move {
let addr_getter = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = addr_getter.local_addr().unwrap();
let (first_sender, first_receiver) = shared_channel::new_bounded(1);
let (second_sender, second_receiver) = shared_channel::new_bounded(1);
let ex1 = LocalExecutorBuilder::new()
.spawn(move || async move {
let receiver = first_receiver.connect().await;
let _ = TcpListener::bind(addr).unwrap();
receiver.recv().await.unwrap();
})
.unwrap();
let ex2 = LocalExecutorBuilder::new()
.spawn(move || async move {
let receiver = second_receiver.connect().await;
let _ = TcpListener::bind(addr).unwrap();
receiver.recv().await.unwrap();
})
.unwrap();
Timer::new(Duration::from_millis(100)).await;
let sender = first_sender.connect().await;
sender.try_send(0).unwrap();
let sender = second_sender.connect().await;
sender.try_send(0).unwrap();
ex1.join().unwrap();
ex2.join().unwrap();
});
}
#[test]
fn multi_executor_accept() {
let (sender, receiver) = shared_channel::new_bounded(1);
let (addr_sender, addr_receiver) = shared_channel::new_bounded(1);
let connected = Arc::new(AtomicUsize::new(0));
let status = connected.clone();
let ex1 = LocalExecutorBuilder::new()
.spawn(move || async move {
let sender = sender.connect().await;
let addr_sender = addr_sender.connect().await;
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
addr_sender.try_send(addr).unwrap();
status.store(1, Ordering::Relaxed);
let accepted = listener.shared_accept().await.unwrap();
sender.try_send(accepted).unwrap();
})
.unwrap();
let status = connected.clone();
let ex2 = LocalExecutorBuilder::new()
.spawn(move || async move {
let receiver = receiver.connect().await;
let accepted = receiver.recv().await.unwrap();
let _ = accepted.bind_to_executor();
status.store(2, Ordering::Relaxed);
})
.unwrap();
let ex3 = LocalExecutorBuilder::new()
.spawn(move || async move {
let receiver = addr_receiver.connect().await;
let addr = receiver.recv().await.unwrap();
TcpStream::connect(addr).await.unwrap()
})
.unwrap();
ex1.join().unwrap();
ex2.join().unwrap();
ex3.join().unwrap();
assert_eq!(connected.load(Ordering::Relaxed), 2);
}
#[test]
fn stream_of_connections() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let coord = Rc::new(Cell::new(0));
let listener_handle: Task<Result<()>> = Task::local(enclose! { (coord) async move {
coord.set(1);
listener.incoming().take(4).try_for_each(|addr| {
addr.map(|_| ())
}).await
}});
while coord.get() != 1 {
Local::later().await;
}
let mut handles = Vec::with_capacity(4);
for _ in 0..4 {
handles.push(Task::local(async move { TcpStream::connect(addr).await }).detach());
}
for handle in handles.drain(..) {
handle.await.unwrap().unwrap();
}
listener_handle.await.unwrap();
let res = TcpStream::connect(addr).await;
assert_eq!(res.is_err(), true)
});
}
#[test]
fn parallel_accept() {
test_executor!(async move {
let listener = Rc::new(TcpListener::bind("127.0.0.1:0").unwrap());
let addr = listener.local_addr().unwrap();
let mut handles = Vec::new();
for _ in 0..128 {
handles.push(
Local::local(enclose! { (listener) async move {
let _accept = listener.accept().await.unwrap();
}})
.detach(),
);
}
Timer::new(Duration::from_millis(100)).await;
for _ in 0..128 {
handles.push(
Local::local(async move {
let _stream = TcpStream::connect(addr).await.unwrap();
})
.detach(),
);
}
for handle in handles {
handle.await.unwrap();
}
});
}
#[test]
fn connect_and_ping_pong() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let coord = Rc::new(Cell::new(0));
let listener_handle = Task::<Result<u8>>::local(enclose! { (coord) async move {
coord.set(1);
let mut stream = listener.accept().await?;
let mut byte = [0u8; 1];
let read = stream.read(&mut byte).await?;
assert_eq!(read, 1);
Ok(byte[0])
}})
.detach();
while coord.get() != 1 {
Local::later().await;
}
let mut stream = TcpStream::connect(addr).await.unwrap();
let byte = [65u8; 1];
let b = stream.write(&byte).await.unwrap();
assert_eq!(b, 1);
assert_eq!(listener_handle.await.unwrap().unwrap(), 65u8);
});
}
#[test]
fn test_read_until() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<usize>>::local(async move {
let mut stream = listener.accept().await?;
let mut buf = Vec::new();
stream.read_until(10, &mut buf).await?;
Ok(buf.len())
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
let vec = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10];
let b = stream.write(&vec).await.unwrap();
assert_eq!(b, 10);
assert_eq!(listener_handle.await.unwrap().unwrap(), 10);
});
}
#[test]
fn test_read_line() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<usize>>::local(async move {
let mut stream = listener.accept().await?;
let mut buf = String::new();
stream.read_line(&mut buf).await?;
Ok(buf.len())
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
let b = stream.write(b"line\n").await.unwrap();
assert_eq!(b, 5);
assert_eq!(listener_handle.await.unwrap().unwrap(), 5);
});
}
#[test]
fn test_lines() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<usize>>::local(async move {
let stream = listener.accept().await?;
Ok(stream.lines().count().await)
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
stream.write(b"line1\nline2\nline3\n").await.unwrap();
stream.write(b"line4\nline5\nline6\n").await.unwrap();
stream.close().await.unwrap();
assert_eq!(listener_handle.await.unwrap().unwrap(), 6);
});
}
#[test]
fn multibuf_fill() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<()>>::local(async move {
let mut stream = listener.accept().await?;
let buf = stream.fill_buf().await?;
assert_eq!(&buf[0..4], b"msg1");
stream.consume(4);
let buf = stream.fill_buf().await?;
assert_eq!(buf, b"msg2");
stream.consume(4);
let buf = stream.fill_buf().await?;
assert_eq!(buf.len(), 0);
Ok(())
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
let b = stream.write(b"msg1").await.unwrap();
assert_eq!(b, 4);
stream.write(b"msg2").await.unwrap();
assert_eq!(b, 4);
stream.close().await.unwrap();
listener_handle.await.unwrap().unwrap();
});
}
#[test]
fn overconsume() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<()>>::local(async move {
let mut stream = listener.accept().await?;
let buf = stream.fill_buf().await?;
assert_eq!(buf.len(), 4);
stream.consume(100);
let buf = stream.fill_buf().await?;
assert_eq!(buf.len(), 0);
Ok(())
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
stream.write(b"msg1").await.unwrap();
stream.close().await.unwrap();
listener_handle.await.unwrap().unwrap();
});
}
#[test]
fn repeated_fill_before_consume() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<()>>::local(async move {
let mut stream = listener.accept().await?;
let buf = stream.fill_buf().await?;
assert_eq!(buf, b"msg1");
let buf = stream.fill_buf().await?;
assert_eq!(buf, b"msg1");
stream.consume(4);
let buf = stream.fill_buf().await?;
assert_eq!(buf.is_empty(), true);
Ok(())
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
stream.write(b"msg1").await.unwrap();
stream.close().await.unwrap();
listener_handle.await.unwrap().unwrap();
});
}
#[test]
fn peek() {
test_executor!(async move {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
let listener_handle = Task::<Result<usize>>::local(async move {
let mut stream = listener.accept().await?;
let mut buf = [0u8; 64];
for _ in 0..10 {
let b = stream.peek(&mut buf).await?;
assert_eq!(b, 4);
assert_eq!(&buf[0..4], b"msg1");
}
stream.read(&mut buf).await?;
stream.peek(&mut buf).await
})
.detach();
let mut stream = TcpStream::connect(addr).await.unwrap();
stream.write(b"msg1").await.unwrap();
stream.close().await.unwrap();
let res = listener_handle.await.unwrap().unwrap();
assert_eq!(res, 0);
});
}
}