#![allow(warnings)]
#![feature(is_some_and)]
#![feature(io_slice_advance)]
use std::mem;
use std::ptr;
use std::rc::Rc;
use std::task::Waker;
use std::str::FromStr;
use std::future::Future;
use std::cell::UnsafeCell;
use std::any::{Any, TypeId};
use std::result::Result as GenResult;
use std::io::{Error, Result, ErrorKind};
use std::sync::{Arc, atomic::{AtomicBool, Ordering}};
use std::net::{SocketAddr, IpAddr, Ipv4Addr, Ipv6Addr};
use std::sync::atomic::AtomicUsize;
use futures::future::LocalBoxFuture;
use crossbeam_channel::Sender;
use mio::{Token, Interest, Poll,
net::TcpStream};
use bytes::BytesMut;
use log::debug;
use pi_async_rt::rt::{serial::AsyncValue,
serial_local_thread::LocalTaskRuntime};
use pi_hash::XHashMap;
#[macro_use]
extern crate lazy_static;
pub mod acceptor;
pub mod connect;
pub mod tls_connect;
pub mod connect_pool;
pub mod server;
pub mod utils;
use utils::{TlsConfig, SocketContext, Hibernate, Ready};
pub const DEFAULT_TCP_IP_V4: &str = "0.0.0.0";
pub const DEFAULT_TCP_IP_V6: &str = "::";
pub const DEFAULT_TCP_PORT: u16 = 38080;
pub const DEFAULT_BUFFER_SIZE: usize = 16384;
pub trait SocketAdapter: Send + Sync + 'static {
type Connect: Socket;
fn connected(&self,
result: GenResult<SocketHandle<Self::Connect>, (SocketHandle<Self::Connect>, Error)>) -> LocalBoxFuture<'static, ()>;
fn readed(&self,
result: GenResult<SocketHandle<Self::Connect>, (SocketHandle<Self::Connect>, Error)>) -> LocalBoxFuture<'static, ()>;
fn writed(&self,
result: GenResult<SocketHandle<Self::Connect>, (SocketHandle<Self::Connect>, Error)>) -> LocalBoxFuture<'static, ()>;
fn closed(&self,
result: GenResult<SocketHandle<Self::Connect>, (SocketHandle<Self::Connect>, Error)>) -> LocalBoxFuture<'static, ()>;
fn timeouted(&self,
handle: SocketHandle<Self::Connect>,
event: SocketEvent) -> LocalBoxFuture<'static, ()>;
}
pub trait SocketAdapterFactory {
type Connect: Socket;
type Adapter: SocketAdapter<Connect = Self::Connect>;
fn get_instance(&self) -> Self::Adapter;
}
#[derive(Clone)]
pub struct SocketOption {
pub recv_buffer_size: usize, pub send_buffer_size: usize, pub read_buffer_capacity: usize, pub write_buffer_capacity: usize, }
impl Default for SocketOption {
fn default() -> Self {
SocketOption {
recv_buffer_size: DEFAULT_BUFFER_SIZE, send_buffer_size: DEFAULT_BUFFER_SIZE, read_buffer_capacity: DEFAULT_BUFFER_SIZE, write_buffer_capacity: 16, }
}
}
#[derive(Clone)]
pub enum SocketConfig {
Raw(Vec<u16>, SocketOption), Tls(Vec<(u16, TlsConfig)>, SocketOption), RawIpv4(IpAddr, Vec<u16>, SocketOption), TlsIpv4(IpAddr, Vec<(u16, TlsConfig)>, SocketOption), RawIpv6(Ipv6Addr, Vec<u16>, SocketOption), TlsIpv6(Ipv6Addr, Vec<(u16, TlsConfig)>, SocketOption), }
impl Default for SocketConfig {
fn default() -> Self {
SocketConfig::new(DEFAULT_TCP_IP_V4, &[DEFAULT_TCP_PORT])
}
}
impl SocketConfig {
pub fn new(ip: &str, port: &[u16]) -> Self {
if ip == DEFAULT_TCP_IP_V6 {
SocketConfig::Raw(port.to_vec(), SocketOption::default())
} else {
let addr: IpAddr;
if let Ok(r) = Ipv4Addr::from_str(ip) {
addr = IpAddr::V4(r);
} else {
if let Ok(r) = Ipv6Addr::from_str(ip) {
addr = IpAddr::V6(r);
} else {
panic!("invalid ip");
}
}
SocketConfig::RawIpv4(addr, port.to_vec(), SocketOption::default())
}
}
pub fn with_tls(ip: &str, port: &[(u16, TlsConfig)]) -> Self {
if ip == DEFAULT_TCP_IP_V6 {
SocketConfig::Tls(port.to_vec(), SocketOption::default())
} else {
let addr: IpAddr;
if let Ok(r) = Ipv4Addr::from_str(ip) {
addr = IpAddr::V4(r);
} else {
if let Ok(r) = Ipv6Addr::from_str(ip) {
addr = IpAddr::V6(r);
} else {
panic!("invalid ip");
}
}
SocketConfig::TlsIpv4(addr, port.to_vec(), SocketOption::default())
}
}
pub fn into_ipv6(self) -> Self {
match self {
SocketConfig::RawIpv4(ip, ports, option) => {
match ip {
IpAddr::V4(addr) => {
SocketConfig::RawIpv6(addr.to_ipv6_mapped(), ports, option)
},
IpAddr::V6(addr) => {
SocketConfig::RawIpv6(addr, ports, option)
},
}
},
SocketConfig::TlsIpv4(ip, ports, option) => {
match ip {
IpAddr::V4(addr) => {
SocketConfig::TlsIpv6(addr.to_ipv6_mapped(), ports, option)
},
IpAddr::V6(addr) => {
SocketConfig::TlsIpv6(addr, ports, option)
},
}
},
config => {
config
}
}
}
pub fn option(&self) -> SocketOption {
match self {
SocketConfig::Raw(_ports, option) => option.clone(),
SocketConfig::Tls(_ports, option) => option.clone(),
SocketConfig::RawIpv4(_ip, _ports, option) => option.clone(),
SocketConfig::TlsIpv4(_ip, _ports, option) => option.clone(),
SocketConfig::RawIpv6(_ip, _ports, option) => option.clone(),
SocketConfig::TlsIpv6(_ip, _ports, option) => option.clone(),
}
}
pub fn set_option(&mut self,
recv_buffer_size: usize,
send_buffer_size: usize,
read_buffer_capacity: usize,
write_buffer_capacity: usize) {
let option = match self {
SocketConfig::Raw(_ports, option) => {
option
},
SocketConfig::Tls(_ports, option) => {
option
},
SocketConfig::RawIpv4(_ip, _ports, option) => {
option
},
SocketConfig::TlsIpv4(_ip, _ports, option) => {
option
},
SocketConfig::RawIpv6(_ip, _ports, option) => {
option
},
SocketConfig::TlsIpv6(_ip, _ports, option) => {
option
},
};
option.recv_buffer_size = recv_buffer_size;
option.send_buffer_size = send_buffer_size;
option.read_buffer_capacity = read_buffer_capacity;
option.write_buffer_capacity = write_buffer_capacity;
}
pub fn addrs(&self) -> Vec<(SocketAddr, TlsConfig)> {
let mut addrs = Vec::with_capacity(1);
match self {
SocketConfig::Raw(ports, _option) => {
for port in ports {
addrs.push((SocketAddr::new(IpAddr::V4(Ipv4Addr::from_str(DEFAULT_TCP_IP_V4).unwrap()), port.clone()), TlsConfig::empty()));
addrs.push((SocketAddr::new(IpAddr::V6(Ipv6Addr::from_str(DEFAULT_TCP_IP_V6).unwrap()), port.clone()), TlsConfig::empty()));
}
},
SocketConfig::Tls(ports, _option) => {
for (port, tls_cfg) in ports {
addrs.push((SocketAddr::new(IpAddr::V4(Ipv4Addr::from_str(DEFAULT_TCP_IP_V4).unwrap()), port.clone()), tls_cfg.clone()));
addrs.push((SocketAddr::new(IpAddr::V6(Ipv6Addr::from_str(DEFAULT_TCP_IP_V6).unwrap()), port.clone()), tls_cfg.clone()));
}
},
SocketConfig::RawIpv4(ip, ports, _option) => {
for port in ports {
addrs.push((SocketAddr::new(ip.clone(), port.clone()), TlsConfig::empty()))
}
},
SocketConfig::TlsIpv4(ip, ports, _option) => {
for (port, tls_cfg) in ports {
addrs.push((SocketAddr::new(ip.clone(), port.clone()), tls_cfg.clone()))
}
},
SocketConfig::RawIpv6(ip, ports, _option) => {
for port in ports {
addrs.push((SocketAddr::new(IpAddr::V6(ip.clone()), port.clone()), TlsConfig::empty()))
}
},
SocketConfig::TlsIpv6(ip, ports, _option) => {
for (port, tls_cfg) in ports {
addrs.push((SocketAddr::new(IpAddr::V6(ip.clone()), port.clone()), tls_cfg.clone()))
}
},
}
addrs
}
pub fn configs(&self) -> Vec<(SocketAddr, SocketOption, TlsConfig)> {
let mut configs = Vec::with_capacity(1);
match self {
SocketConfig::Raw(ports, option) => {
for port in ports {
configs.push((SocketAddr::new(IpAddr::V4(Ipv4Addr::from_str(DEFAULT_TCP_IP_V4).unwrap()), port.clone()), option.clone(), TlsConfig::empty()));
configs.push((SocketAddr::new(IpAddr::V6(Ipv6Addr::from_str(DEFAULT_TCP_IP_V6).unwrap()), port.clone()), option.clone(), TlsConfig::empty()));
}
},
SocketConfig::Tls(ports, option) => {
for (port, tls_cfg) in ports {
configs.push((SocketAddr::new(IpAddr::V4(Ipv4Addr::from_str(DEFAULT_TCP_IP_V4).unwrap()), port.clone()), option.clone(), tls_cfg.clone()));
configs.push((SocketAddr::new(IpAddr::V6(Ipv6Addr::from_str(DEFAULT_TCP_IP_V6).unwrap()), port.clone()), option.clone(), tls_cfg.clone()));
}
},
SocketConfig::RawIpv4(ip, ports, option) => {
for port in ports {
configs.push((SocketAddr::new(ip.clone(), port.clone()), option.clone(), TlsConfig::empty()))
}
},
SocketConfig::TlsIpv4(ip, ports, option) => {
for (port, tls_cfg) in ports {
configs.push((SocketAddr::new(ip.clone(), port.clone()), option.clone(), tls_cfg.clone()))
}
},
SocketConfig::RawIpv6(ip, ports, option) => {
for port in ports {
configs.push((SocketAddr::new(IpAddr::V6(ip.clone()), port.clone()), option.clone(), TlsConfig::empty()))
}
},
SocketConfig::TlsIpv6(ip, ports, option) => {
for (port, tls_cfg) in ports {
configs.push((SocketAddr::new(IpAddr::V6(ip.clone()), port.clone()), option.clone(), tls_cfg.clone()))
}
},
}
configs
}
}
#[derive(Debug)]
pub enum SocketStatus {
Connected(Result<()>), Readed(Result<()>), Writed(Result<()>), Closed(Result<()>), Timeout(SocketEvent), }
#[derive(Debug)]
pub struct SocketEvent {
inner: Box<dyn Any + Send + 'static>, }
unsafe impl Send for SocketEvent {}
impl Default for SocketEvent {
fn default() -> Self {
let inner = Box::new(());
Self { inner }
}
}
impl SocketEvent {
pub fn empty() -> Self {
Self::default()
}
pub fn is_empty(&self) -> bool {
self.inner.type_id() == TypeId::of::<()>()
}
pub fn get<T: 'static>(&self) -> Option<&T> {
if self.is_empty() {
return None;
}
self.inner.downcast_ref::<T>()
}
pub fn get_mut<T: 'static>(&mut self) -> Option<&mut T> {
if self.is_empty() {
return None;
}
self.inner.downcast_mut::<T>()
}
pub fn set<T: Send + 'static>(&mut self, event: T) -> bool {
if !self.is_empty() {
return false;
}
self.inner = Box::new(event);
true
}
pub fn remove<T: 'static>(&mut self) -> Option<T> {
if self.is_empty() {
return None;
}
let old = mem::replace(&mut self.inner, Box::new(()));
match old.downcast() {
Err(_) => None,
Ok(inner) => Some(*inner),
}
}
}
pub trait Stream: Sized + 'static {
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;
fn set_runtime(&mut self, rt: LocalTaskRuntime<()>);
fn set_handle(&mut self, shared: &Arc<UnsafeCell<Self>>);
fn get_stream_ref(&self) -> &TcpStream;
fn get_stream_mut(&mut self) -> &mut TcpStream;
fn set_token(&mut self, token: Option<Token>) -> Option<Token>;
fn set_uid(&mut self, uid: usize) -> Option<usize>;
fn get_interest(&self) -> Option<Interest>;
fn set_interest(&self, ready: Interest);
fn get_read_block_len(&self) -> usize;
fn set_read_block_len(&self, len: usize);
fn get_write_block_len(&self) -> usize;
fn set_write_block_len(&self, len: usize);
fn set_poll(&mut self, poll: Rc<UnsafeCell<Poll>>);
fn set_write_listener(&mut self, listener: Option<Sender<(Token, Vec<u8>)>>);
fn set_close_listener(&mut self, listener: Option<Sender<(Token, Result<()>)>>);
fn set_timer_listener(&mut self,
listener: Option<Sender<(Token, Option<(usize, SocketEvent)>)>>);
fn set_timer_handle(&mut self, timer: usize) -> Option<usize>;
fn unset_timer_handle(&mut self) -> Option<usize>;
fn is_require_recv(&self) -> bool;
fn recv(&mut self) -> Result<usize>;
fn send(&mut self) -> Result<usize>;
}
pub trait Socket: Sized + 'static {
fn is_closed(&self) -> bool;
fn is_flush(&self) -> bool;
fn set_flush(&self, flush: bool);
fn get_handle(&self) -> SocketHandle<Self>;
fn remove_handle(&mut self) -> Option<SocketHandle<Self>>;
fn get_local(&self) -> &SocketAddr;
fn get_remote(&self) -> &SocketAddr;
fn get_token(&self) -> Option<&Token>;
fn get_uid(&self) -> Option<&usize>;
fn get_context(&self) -> Rc<UnsafeCell<SocketContext>>;
fn set_timeout(&self, timeout: usize, event: SocketEvent);
fn unset_timeout(&self);
fn is_security(&self) -> bool;
fn read_ready(&mut self, adjust: usize) -> GenResult<AsyncValue<usize>, usize>;
fn is_wait_wakeup_read_ready(&self) -> bool;
fn wakeup_read_ready(&mut self);
fn get_read_buffer(&self) -> Rc<UnsafeCell<Option<BytesMut>>>;
fn get_write_buffer(&mut self) -> Option<&mut BytesMut>;
fn write_ready<B>(&mut self, buf: B) -> Result<()>
where B: AsRef<[u8]> + 'static;
fn reregister_interest(&mut self, ready: Ready) -> Result<()>;
fn is_hibernated(&self) -> bool;
fn push_hibernated_task<F>(&self,
task: F)
where F: Future<Output = ()> + 'static;
fn run_hibernated_tasks(&self);
fn hibernate(&self,
handle: SocketHandle<Self>,
ready: Ready) -> Option<Hibernate<Self>>;
fn set_hibernate(&self, hibernate: Hibernate<Self>) -> bool;
fn set_hibernate_wakers(&self, waker: Waker);
fn wakeup(&mut self, result: Result<()>) -> bool;
fn close(&mut self, reason: Result<()>) -> Result<()>;
}
#[derive(Clone)]
pub enum AcceptorCmd {
Continue, Pause(usize), Close(String), }
pub struct SocketDriver<S: Socket + Stream, A: SocketAdapter<Connect = S>> {
addrs: Rc<XHashMap<SocketAddr, usize>>, controller: Option<Rc<Sender<Box<dyn FnOnce() -> AcceptorCmd + Send>>>>, count: Rc<AtomicUsize>, router: Rc<Vec<Sender<S>>>, adapter: Option<Arc<A>>, }
unsafe impl<S: Socket + Stream, A: SocketAdapter<Connect = S>> Send for SocketDriver<S, A> {}
unsafe impl<S: Socket + Stream, A: SocketAdapter<Connect = S>> Sync for SocketDriver<S, A> {}
impl<S: Socket + Stream, A: SocketAdapter<Connect = S>> Clone for SocketDriver<S, A> {
fn clone(&self) -> Self {
SocketDriver {
addrs: self.addrs.clone(),
controller: self.controller.clone(),
count: self.count.clone(),
router: self.router.clone(),
adapter: self.adapter.clone(),
}
}
}
impl<S: Socket + Stream, A: SocketAdapter<Connect = S>> SocketDriver<S, A> {
pub fn new(bind: &[(SocketAddr, Sender<S>)]) -> Self {
let size = bind.len();
let mut map = XHashMap::default();
let mut vec = Vec::with_capacity(size);
let mut index: usize = 0;
for (addr, sender) in bind {
map.insert(addr.clone(), index);
vec.push(sender.clone());
index += 1;
}
let count = AtomicUsize::new(0);
SocketDriver {
addrs: Rc::new(map),
controller: None,
count: Rc::new(count),
router: Rc::new(vec),
adapter: None,
}
}
pub fn get_addrs(&self) -> Vec<SocketAddr> {
self.addrs.keys().map(|addr| {
addr.clone()
}).collect::<Vec<SocketAddr>>()
}
pub fn get_controller(&self) -> Option<&Rc<Sender<Box<dyn FnOnce() -> AcceptorCmd + Send>>>> {
self.controller.as_ref()
}
pub fn set_controller(&mut self, controller: Sender<Box<dyn FnOnce() -> AcceptorCmd + Send>>) {
self.controller = Some(Rc::new(controller));
}
pub fn route(&self, mut socket: S) -> Result<()> {
if let Some(_token) = socket.set_token(None) {
let router = &self.router[self.count.fetch_add(1, Ordering::Relaxed) % self.router.len()];
match router.try_send(socket) {
Err(e) => {
Err(Error::new(ErrorKind::BrokenPipe,
format!("tcp socket route failed, e: {:?}",
e)))
},
Ok(_) => Ok(()),
}
} else {
Err(Error::new(ErrorKind::Interrupted,
format!("tcp socket route failed, e: invalid accept token")))
}
}
pub fn get_adapter(&self) -> &A {
self.adapter.as_ref().unwrap()
}
pub fn clone_adapter(&self) -> Arc<A> {
self
.adapter
.as_ref()
.unwrap()
.clone()
}
pub fn set_adapter(&mut self, adapter: A) {
self.adapter = Some(Arc::new(adapter));
}
}
pub struct SocketHandle<S: Socket>(Arc<SocketImage<S>>);
unsafe impl<S: Socket> Send for SocketHandle<S> {}
unsafe impl<S: Socket> Sync for SocketHandle<S> {}
impl<S: Socket> Clone for SocketHandle<S> {
fn clone(&self) -> Self {
SocketHandle(self.0.clone())
}
}
impl<S: Socket> Drop for SocketHandle<S> {
fn drop(&mut self) {
debug!("Drop socket handle, token: {:?}, uid: {:?}, remote: {:?}, local: {:?}, closed: {:?}, socket image shared: {:?}",
self.0.token,
self.0.uid,
self.0.remote,
self.0.local,
self.0.closed.load(Ordering::Relaxed),
Arc::strong_count(&self.0));
}
}
impl<S: Socket> SocketHandle<S> {
pub fn new(image: SocketImage<S>) -> Self {
SocketHandle(Arc::new(image))
}
pub fn is_closed(&self) -> bool {
self.0.closed.load(Ordering::Acquire)
}
pub fn get_token(&self) -> &Token {
&self.0.token
}
pub fn get_uid(&self) -> usize {
self.0.uid
}
pub fn get_local(&self) -> &SocketAddr {
&self.0.local
}
pub fn get_remote(&self) -> &SocketAddr {
&self.0.remote
}
pub fn set_timeout(&self, timeout: usize, event: SocketEvent) {
let _ = self
.0
.timer_listener
.send((self.0.token, Some((timeout, event))));
}
pub fn unset_timeout(&self) {
let _ = self
.0
.timer_listener
.send((self.0.token, None));
}
pub fn is_security(&self) -> bool {
self.0.security
}
pub fn get_context(&self) -> Rc<UnsafeCell<SocketContext>> {
unsafe {
(&*self.0.inner.get()).get_context()
}
}
pub fn read_ready(&self, size: usize) -> GenResult<AsyncValue<usize>, usize> {
unsafe {
(&mut *self.0.inner.get()).read_ready(size)
}
}
pub fn get_read_buffer(&self) -> Rc<UnsafeCell<Option<BytesMut>>> {
unsafe {
(&*self.0.inner.get()).get_read_buffer()
}
}
pub fn write_ready<B>(&self, buf: B) -> Result<()>
where B: AsRef<[u8]> + 'static {
unsafe {
(&mut *self.0.inner.get()).write_ready(buf)
}
}
pub fn reregister_interest(&self, ready: Ready) -> Result<()> {
unsafe {
(&mut *self.0.inner.get()).reregister_interest(ready)
}
}
pub fn run_hibernated_tasks(&self) {
unsafe {
(&*self.0.inner.get()).run_hibernated_tasks();
}
}
pub fn hibernate(&self,
handle: SocketHandle<S>,
ready: Ready) -> Option<Hibernate<S>> {
unsafe {
(&*self.0.inner.get()).hibernate(handle, ready)
}
}
pub fn set_hibernate(&self, hibernate: Hibernate<S>) -> bool {
unsafe {
(&*self.0.inner.get()).set_hibernate(hibernate)
}
}
fn set_hibernate_wakers(&self, waker: Waker) {
unsafe {
(&*self.0.inner.get()).set_hibernate_wakers(waker);
}
}
pub fn wakeup(&self, result: Result<()>) -> bool {
unsafe {
(&mut *self.0.inner.get()).wakeup(result)
}
}
pub fn close(&self, reason: Result<()>) -> Result<()> {
unsafe {
(&mut *self.0.inner.get()).close(reason)
}
}
}
pub struct SocketImage<S: Socket> {
inner: Arc<UnsafeCell<S>>, local: SocketAddr, remote: SocketAddr, token: Token, uid: usize, security: bool, closed: Arc<AtomicBool>, close_listener: Sender<(Token, Result<()>)>, timer_listener: Sender<(Token, Option<(usize, SocketEvent)>)>, }
unsafe impl<S: Socket> Send for SocketImage<S> {}
unsafe impl<S: Socket> Sync for SocketImage<S> {}
impl<S: Socket> Drop for SocketImage<S> {
fn drop(&mut self) {
debug!("Drop socket image, token: {:?}, uid: {:?}, remote: {:?}, local: {:?}, closed: {:?}, socket shared: {:?}",
self.token,
self.uid,
self.remote,
self.local,
self.closed.load(Ordering::Relaxed),
Arc::strong_count(&self.inner));
}
}
impl<S: Socket> SocketImage<S> {
pub fn new(shared: &Arc<UnsafeCell<S>>,
local: SocketAddr,
remote: SocketAddr,
token: Token,
uid: usize,
security: bool,
closed: Arc<AtomicBool>,
close_listener: Sender<(Token, Result<()>)>,
timer_listener: Sender<(Token, Option<(usize, SocketEvent)>)>
) -> Self {
SocketImage {
inner: shared.clone(),
local,
remote,
token,
uid,
security,
closed,
close_listener,
timer_listener,
}
}
}
pub trait AsyncService<S: Socket>: Send + Sync + 'static {
fn handle_connected(&self, handle: SocketHandle<S>, status: SocketStatus) -> LocalBoxFuture<'static, ()>;
fn handle_readed(&self, handle: SocketHandle<S>, status: SocketStatus) -> LocalBoxFuture<'static, ()>;
fn handle_writed(&self, handle: SocketHandle<S>, status: SocketStatus) -> LocalBoxFuture<'static, ()>;
fn handle_closed(&self, handle: SocketHandle<S>, status: SocketStatus) -> LocalBoxFuture<'static, ()>;
fn handle_timeouted(&self, handle: SocketHandle<S>, status: SocketStatus) -> LocalBoxFuture<'static, ()>;
}