use std::rc::Rc;
use std::task::Waker;
use std::future::Future;
use std::borrow::Borrow;
use std::net::SocketAddr;
use std::cell::UnsafeCell;
use std::collections::VecDeque;
use std::result::Result as GenResult;
use std::io::{Error, Result, ErrorKind, Read, Write};
use std::sync::{Arc, atomic::{AtomicBool, AtomicUsize, Ordering}};
use mio::{Token, Interest, Poll,
net::TcpStream};
use crossbeam_channel::Sender;
use futures::{sink::SinkExt,
future::{FutureExt, LocalBoxFuture}};
use bytes::{Buf, BufMut, BytesMut};
use log::debug;
use pi_async_rt::{lock::spin_lock::SpinLock,
rt::{serial::AsyncValue,
serial_local_thread::LocalTaskRuntime}};
use pi_async_buffer::async_pipeline::{AsyncReceiverExt, AsyncPipeLineExt, PipeSender, channel};
use crate::{Stream, Socket, SocketEvent, SocketHandle, SocketImage, SocketContext,
utils::{TlsConfig, Hibernate, Ready}};
const DEFAULT_READ_BLOCK_LEN: usize = 4096;
const DEFAULT_WRITE_BLOCK_LEN: usize = 4096;
const DEFAULT_READ_BLOCK_ADJUST_LEN: usize = 1024;
const DEAFULT_RECV_FRAME_BUF_SIZE: usize = 16;
const MIN_READED_READ_BUF_SIZE_LIMIT: usize = DEFAULT_READ_BLOCK_LEN;
const MIN_READED_WRITE_BUF_SIZE_LIMIT: usize = DEFAULT_WRITE_BLOCK_LEN;
const DEAFULT_READED_READ_BUF_SIZE_LIMIT: usize = 256 * 1024;
const DEAFULT_READED_WRITE_BUF_SIZE_LIMIT: usize = 256 * 1024;
pub struct TcpSocket {
rt: Option<LocalTaskRuntime<()>>,
uid: Option<usize>,
local: SocketAddr,
remote: SocketAddr,
token: Option<Token>,
stream: TcpStream,
interest: Arc<SpinLock<Interest>>,
wait_recv_len: usize,
recv_len: usize,
read_len: Arc<AtomicUsize>,
readed_read_limit: Arc<AtomicUsize>,
readed: usize,
read_buf: Rc<UnsafeCell<Option<BytesMut>>>,
wait_ready_len: usize,
ready_len: usize,
ready_reader: SpinLock<Option<AsyncValue<usize>>>,
wait_sent_len: AtomicUsize,
sent_len: usize,
write_len: Arc<AtomicUsize>,
readed_write_limit: Arc<AtomicUsize>,
readed_write_len: usize,
write_buf: Option<BytesMut>,
poll: Option<Rc<UnsafeCell<Poll>>>,
hibernate: SpinLock<Option<Hibernate<Self>>>,
hibernate_wakers: SpinLock<VecDeque<Waker>>,
hibernated_queue: Arc<SpinLock<VecDeque<LocalBoxFuture<'static, ()>>>>,
handle: Option<SocketHandle<Self>>,
context: Rc<UnsafeCell<SocketContext>>,
write_listener: Option<Sender<(Token, Vec<u8>)>>,
closed: Arc<AtomicBool>,
close_listener: Option<Sender<(Token, Result<()>)>>,
timer_handle: Option<usize>,
timer_listener: Option<Sender<(Token, Option<(usize, SocketEvent)>)>>,
}
unsafe impl Send for TcpSocket {}
unsafe impl Sync for TcpSocket {}
impl Drop for TcpSocket {
fn drop(&mut self) {
debug!("Drop tcp socket, token: {:?}, uid: {:?}, remote: {:?}, local: {:?}, closed: {:?}",
self.token,
self.uid,
self.remote,
self.local,
self.closed.load(Ordering::Relaxed));
}
}
impl Stream for TcpSocket {
fn new(local: &SocketAddr,
remote: &SocketAddr,
token: Option<Token>,
stream: TcpStream,
recv_frame_buf_size: usize,
readed_read_size_limit: usize,
readed_write_size_limit: usize,
_tls_cfg: TlsConfig) -> Self {
let recv_frame_buf_size = if recv_frame_buf_size == 0 {
DEAFULT_RECV_FRAME_BUF_SIZE
} else {
recv_frame_buf_size
};
let readed_read_size_limit = if readed_read_size_limit < MIN_READED_READ_BUF_SIZE_LIMIT {
DEAFULT_READED_READ_BUF_SIZE_LIMIT
} else {
readed_read_size_limit
};
let readed_write_size_limit = if readed_write_size_limit < MIN_READED_WRITE_BUF_SIZE_LIMIT {
DEAFULT_READED_WRITE_BUF_SIZE_LIMIT
} else {
readed_write_size_limit
};
let interest = Arc::new(SpinLock::new(Interest::READABLE));
let read_len = Arc::new(AtomicUsize::new(DEFAULT_READ_BLOCK_LEN));
let readed_read_limit = Arc::new(AtomicUsize::new(readed_read_size_limit));
let read_buf = Rc::new(UnsafeCell::new(Some(BytesMut::new())));
let ready_reader = SpinLock::new(None);
let wait_sent_len = AtomicUsize::new(0);
let write_len = Arc::new(AtomicUsize::new(DEFAULT_WRITE_BLOCK_LEN));
let readed_write_limit = Arc::new(AtomicUsize::new(readed_write_size_limit));
let write_buf = Some(BytesMut::new());
let hibernate = SpinLock::new(None);
let hibernate_wakers = SpinLock::new(VecDeque::new());
let hibernated_queue = Arc::new(SpinLock::new(VecDeque::new()));
let context = Rc::new(UnsafeCell::new(SocketContext::empty()));
let closed = Arc::new(AtomicBool::new(false));
TcpSocket {
rt: None,
uid: None,
local: local.clone(),
remote: remote.clone(),
token,
stream,
interest,
wait_recv_len: 0,
recv_len: 0,
read_len,
readed_read_limit,
readed: 0,
read_buf,
wait_ready_len: 0,
ready_len: 0,
ready_reader,
wait_sent_len,
sent_len: 0,
write_len,
readed_write_limit,
readed_write_len: 0,
write_buf,
poll: None,
hibernate,
hibernate_wakers,
hibernated_queue,
handle: None,
context,
write_listener: None,
closed,
close_listener: None,
timer_handle: None,
timer_listener: None,
}
}
fn set_runtime(&mut self, rt: LocalTaskRuntime<()>) {
self.rt = Some(rt);
}
fn set_handle(&mut self, shared: &Arc<UnsafeCell<Self>>) {
if let Some(token) = self.token {
if let Some(close_listener) = &self.close_listener {
if let Some(timer_listener) = &self.timer_listener {
let image = SocketImage::new(shared,
self.local,
self.remote,
token,
self.uid.unwrap(),
self.is_security(),
self.closed.clone(),
close_listener.clone(),
timer_listener.clone()
);
self.handle = Some(SocketHandle::new(image));
}
}
}
}
#[inline]
fn get_stream_ref(&self) -> &TcpStream {
&self.stream
}
#[inline]
fn get_stream_mut(&mut self) -> &mut TcpStream {
&mut self.stream
}
fn set_token(&mut self, token: Option<Token>) -> Option<Token> {
let last = self.token.take();
self.token = token;
last
}
fn set_uid(&mut self, uid: usize) -> Option<usize> {
let last = self.uid.take();
self.uid = Some(uid);
last
}
fn get_interest(&self) -> Option<Interest> {
Some({ *self.interest.lock() })
}
fn set_interest(&self, interest: Interest) {
*self.interest.lock() = interest;
}
#[inline]
fn get_read_block_len(&self) -> usize {
self.read_len.load(Ordering::Acquire)
}
fn set_read_block_len(&self, len: usize) {
self.read_len.store(len, Ordering::Release);
}
fn get_write_block_len(&self) -> usize {
self.write_len.load(Ordering::Acquire)
}
fn set_write_block_len(&self, len: usize) {
self.write_len.store(len, Ordering::Release);
}
fn set_poll(&mut self, poll: Rc<UnsafeCell<Poll>>) {
self.poll = Some(poll);
}
fn set_write_listener(&mut self,
listener: Option<Sender<(Token, Vec<u8>)>>) {
self.write_listener = listener;
}
fn set_close_listener(&mut self,
listener: Option<Sender<(Token, Result<()>)>>) {
self.close_listener = listener;
}
fn set_timer_listener(&mut self,
listener: Option<Sender<(Token, Option<(usize, SocketEvent)>)>>) {
self.timer_listener = listener;
}
fn set_timer_handle(&mut self, timer_handle: usize) -> Option<usize> {
let last_timer_handle = self.unset_timer_handle();
self.timer_handle = Some(timer_handle);
last_timer_handle
}
fn unset_timer_handle(&mut self) -> Option<usize> {
self.timer_handle.take()
}
#[inline]
fn is_require_recv(&self) -> bool {
self.wait_recv_len > self.recv_len
}
fn recv(&mut self) -> Result<usize> {
if self.is_closed() {
let token = self.get_token().unwrap().clone();
let remote = self.get_remote().clone();
let local = self.get_local().clone();
return Err(Error::new(ErrorKind::ConnectionAborted,
format!("Receive stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: connection already closed",
token,
remote,
local)));
}
let mut block_pos = 0; let mut block = Vec::with_capacity(self.get_read_block_len()); block.resize(self.get_read_block_len(), 0);
let mut result = Ok(0); loop {
match self.get_stream_mut().read(&mut block[block_pos..]) {
Ok(0) => {
result = Err(Error::new(ErrorKind::ConnectionAborted,
format!("Receive stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: peer already closed",
self.get_token(),
self.get_remote(),
self.get_local())));
break;
},
Ok(len) => {
self.recv_len += len; if self.ready_reader.lock().is_some() {
self.ready_len += len;
}
block_pos += len;
if block_pos == block.len() {
block.resize(block.len() + DEFAULT_READ_BLOCK_ADJUST_LEN, 0);
}
},
Err(ref e) if e.kind() == ErrorKind::Interrupted => {
continue;
},
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
block.truncate(block_pos); if let Some(buf) = unsafe { (&mut *self.read_buf.get()) } {
buf.put_slice(&block[..]);
}
result = Ok(block_pos);
break;
},
Err(e) => {
result = Err(Error::new(e.kind(),
format!("Receive stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: {:?}",
self.get_token(),
self.get_remote(),
self.get_local(),
e)));
break;
},
}
}
if !self.is_require_recv() {
self.set_interest(Interest::WRITABLE); }
if result.is_ok() {
if self.readed > self.readed_read_limit.load(Ordering::Relaxed) {
unsafe {
let old_buf = (&mut *self.read_buf.get()).take().unwrap();
let mut new_buf = BytesMut::with_capacity(old_buf.remaining());
new_buf.put(old_buf);
*self.read_buf.get() = Some(new_buf);
self.readed = 0;
}
}
}
result
}
fn send(&mut self) -> Result<usize> {
if self.is_closed() {
return Err(Error::new(ErrorKind::ConnectionAborted,
format!("Send stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: connection already closed",
self.get_token(),
self.get_remote(),
self.get_local())));
}
let mut result = Ok(0); let mut block_pos = 0; let mut block_len = self.get_write_block_len();
if let Some(write_buf) = self.get_write_buffer() {
let remaining = write_buf.remaining(); block_len = if block_len > remaining {
remaining
} else {
block_len
};
let mut bytes = write_buf.copy_to_bytes(block_len);
let mut block = bytes.as_ref(); drop(write_buf);
while !block.is_empty() {
match self.get_stream_mut().write(block) {
Ok(0) => {
result = Err(Error::new(ErrorKind::WriteZero,
format!("Send stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: failed to write whole buffer",
self.get_token(),
self.get_remote(),
self.get_local())));
break;
},
Ok(len) => {
self.sent_len += len; self.readed_write_len += len; block_pos += len;
block = &block[block_pos..]; result = Ok(block_pos);
},
Err(e) if e.kind() == ErrorKind::Interrupted => {
continue;
},
Err(e) if e.kind() == ErrorKind::WouldBlock => {
if !block.is_empty() {
if let Some(write_buf) = &mut self.write_buf {
write_buf.put_slice(&block[block_pos..]);
}
}
result = Ok(block_pos);
break;
},
Err(e) => {
result = Err(Error::new(e.kind(),
format!("Send stream failed, token: {:?}, peer: {:?}, local: {:?}, reason: {:?}",
self.get_token(),
self.get_remote(),
self.get_local(),
e)));
break;
},
}
}
}
if (block_pos > 0)
&& (self.readed_write_len > self.readed_write_limit.load(Ordering::Relaxed)) {
let old_write_buf = self
.write_buf
.take()
.unwrap();
let mut new_write_buf = BytesMut::new();
new_write_buf.put(old_write_buf);
self.readed_write_len = 0; self.write_buf = Some(new_write_buf); }
if block_pos < block_len {
let interest = { *self.interest.lock() }; self.set_interest(interest.add(Interest::WRITABLE));
} else {
if let Some(write_buf) = &mut self.write_buf {
if write_buf.remaining() > 0
|| self.wait_sent_len.load(Ordering::Relaxed) > self.sent_len {
self.set_interest(Interest::WRITABLE);
} else {
self.set_interest(Interest::READABLE);
}
} else {
self.set_interest(Interest::READABLE);
}
}
result
}
}
impl Socket for TcpSocket {
fn is_closed(&self) -> bool {
self.closed.load(Ordering::Acquire)
}
fn is_flush(&self) -> bool {
true
}
fn set_flush(&self, flush: bool) {
}
fn get_handle(&self) -> SocketHandle<Self> {
self
.handle
.as_ref()
.unwrap()
.clone()
}
fn remove_handle(&mut self) -> Option<SocketHandle<Self>> {
self.handle.take()
}
fn get_local(&self) -> &SocketAddr {
&self.local
}
fn get_remote(&self) -> &SocketAddr {
&self.remote
}
fn get_token(&self) -> Option<&Token> {
self.token.as_ref()
}
fn get_uid(&self) -> Option<&usize> {
self.uid.as_ref()
}
fn get_context(&self) -> Rc<UnsafeCell<SocketContext>> {
self.context.clone()
}
fn set_timeout(&self, timeout: usize, event: SocketEvent) {
if let Some(listener) = &self.timer_listener {
if let Some(token) = self.token {
listener.send((token, Some((timeout, event))));
}
}
}
fn unset_timeout(&self) {
if let Some(listener) = &self.timer_listener {
if let Some(token) = self.token {
listener.send((token, None));
}
}
}
fn is_security(&self) -> bool {
false
}
fn read_ready(&mut self, adjust: usize) -> GenResult<AsyncValue<usize>, usize> {
if self.is_closed() {
return Err(0);
}
self.wait_recv_len += adjust; let interest = { *self.interest.lock() }; self.set_interest(interest.add(Interest::READABLE));
let remaining = unsafe {
(&*self.read_buf.get())
.as_ref()
.unwrap()
.remaining()
};
if remaining >= adjust && remaining > 0 {
return Err(remaining);
}
let value = AsyncValue::new();
let value_copy = value.clone();
*self.ready_reader.lock() = Some(value); self.wait_ready_len = adjust - remaining;
Ok(value_copy)
}
#[inline]
fn is_wait_wakeup_read_ready(&self) -> bool {
self.ready_reader.lock().is_some()
}
fn wakeup_read_ready(&mut self) {
if (self.wait_ready_len == 0) || (self.wait_ready_len <= self.ready_len) {
if let Some(ready_reader) = self.ready_reader.lock().take() {
ready_reader.set(self.ready_len); self.wait_ready_len = 0; self.ready_len = 0; }
}
}
fn get_read_buffer(&self) -> Rc<UnsafeCell<Option<BytesMut>>> {
self.read_buf.clone()
}
#[inline]
fn get_write_buffer(&mut self) -> Option<&mut BytesMut> {
self.write_buf.as_mut()
}
fn write_ready<B>(&mut self, buf: B) -> Result<()>
where B: AsRef<[u8]> + 'static {
if self.is_closed() {
return Err(Error::new(ErrorKind::ConnectionAborted,
format!("Write ready failed, token: {:?}, peer: {:?}, local: {:?}, reason: connection already closed",
self.get_token(),
self.get_remote(),
self.get_local())));
}
self.wait_sent_len
.fetch_add(buf.as_ref().len(), Ordering::Relaxed); if let Some(listener) = &self.write_listener {
if let Some(token) = self.token.clone() {
listener.send((token, buf.as_ref().to_vec()));
}
}
Ok(())
}
fn reregister_interest(&mut self, ready: Ready) -> Result<()> {
let interest = { *self.interest.lock() }; match ready {
Ready::Empty => self.set_interest(Interest::WRITABLE), Ready::Readable => self.set_interest(interest.add(Interest::READABLE)), Ready::Writable => self.set_interest(interest.add(Interest::WRITABLE)), Ready::OnlyRead => self.set_interest(Interest::READABLE), Ready::OnlyWrite => self.set_interest(Interest::WRITABLE), Ready::ReadWrite => self.set_interest(Interest::READABLE.add(Interest::WRITABLE)) }
if let Some(interest) = self.get_interest() {
let token = self.get_token().unwrap().clone();
unsafe {
(&mut *self
.poll
.as_ref()
.unwrap()
.get())
.registry()
.reregister(self.get_stream_mut(),
token,
interest)
}
} else {
Ok(())
}
}
fn is_hibernated(&self) -> bool {
self.hibernate.lock().is_some() ||
self.hibernated_queue.lock().len() > 0
}
fn push_hibernated_task<F>(&self,
task: F)
where F: Future<Output = ()> + 'static {
let boxed = async move {
task.await;
}.boxed_local();
self.hibernated_queue.lock().push_back(boxed);
}
fn run_hibernated_tasks(&self) {
if let Some(rt) = &self.rt {
let hibernated_queue = self.hibernated_queue.clone();
rt.send(async move {
loop {
let task = {
hibernated_queue
.lock()
.pop_front()
};
if let Some(task) = task {
task.await;
} else {
return;
}
}
});
}
}
fn hibernate(&self,
handle: SocketHandle<Self>,
ready: Ready) -> Option<Hibernate<Self>> {
if self.is_closed() {
return None;
}
let hibernate = Hibernate::new(handle, ready);
let hibernate_copy = hibernate.clone();
Some(hibernate_copy)
}
fn set_hibernate(&self, hibernate: Hibernate<Self>) -> bool {
let mut locked = self.hibernate.lock();
if locked.is_some() {
return false;
}
*locked = Some(hibernate);
true
}
fn set_hibernate_wakers(&self, waker: Waker) {
self
.hibernate_wakers
.lock()
.push_back(waker);
}
fn wakeup(&mut self, result: Result<()>) -> bool {
if self.is_closed() {
return true;
}
let mut r = false;
if let Some(hibernate) = self.hibernate.lock().take() {
r = hibernate.wakeup(result); }
if r {
let mut locked = self.hibernate_wakers.lock();
while let Some(waker) = locked.pop_front() {
waker.wake();
}
}
r
}
fn close(&mut self, reason: Result<()>) -> Result<()> {
if let Ok(true) = self.closed.compare_exchange(false,
true,
Ordering::AcqRel,
Ordering::Relaxed) {
return Ok(());
}
if let Some(value) = self.ready_reader.lock().take() {
value.set(0);
self.wait_ready_len = 0; self.ready_len = 0; }
if let Some(listener) = &self.close_listener {
if let Some(token) = self.get_token() {
if let Err(e) = listener.send((token.clone(), reason)) {
return Err(Error::new(ErrorKind::BrokenPipe, e));
}
}
}
Ok(())
}
}