use std::ptr;
use std::fs::File;
use std::pin::Pin;
use std::path::Path;
use std::result::Result;
use std::future::Future;
use std::ops::{Deref, DerefMut};
use std::cell::{RefCell, Ref, RefMut};
use std::task::{Poll, Context, Waker};
use std::sync::{Arc, atomic::{AtomicU8, Ordering}};
use std::io::{Read, BufReader, Result as IOResult, Error, ErrorKind};
use crossbeam_channel::Sender;
use parking_lot::RwLock;
use mio::Token;
use rustls::{ALL_CIPHER_SUITES, ProtocolVersion, RootCertStore, Certificate, PrivateKey, ClientConfig, ServerConfig,
server::{NoClientAuth,
AllowAnyAuthenticatedClient,
AllowAnyAnonymousOrAuthenticatedClient}};
use rustls_pemfile;
use pi_async_rt::{lock::spin_lock::SpinLock,
rt::{serial::AsyncWaitResult,
serial_local_thread::LocalTaskRuntime}};
use pi_hash::XHashMap;
use crate::{Socket, SocketHandle, Stream};
lazy_static! {
pub static ref TCP_SOCKET_POOL_SENDER_TAB: Arc<RwLock<XHashMap<u8, Sender<(Token, IOResult<()>)>>>> = Arc::new(RwLock::new(XHashMap::default()));
}
pub fn register_close_sender(uid: u8, sender: Sender<(Token, IOResult<()>)>) {
TCP_SOCKET_POOL_SENDER_TAB.write().insert(uid, sender);
}
pub fn close_socket(uid: usize, reason: IOResult<()>) -> bool {
let pool_uid = (uid >> 24 & 0xff) as u8;
let token = Token(uid & 0xffffff);
if let Some(sender) = TCP_SOCKET_POOL_SENDER_TAB.read().get(&pool_uid) {
sender.send((token, reason));
return true;
}
false
}
#[derive(Clone)]
pub enum TlsConfig {
Empty, Server(Arc<ServerConfig>), Client(Arc<ClientConfig>), }
impl TlsConfig {
pub fn empty() -> Self {
TlsConfig::Empty
}
pub fn new_server(client_auth_path: &str, is_client_auth: bool, server_certs_path: &str, server_key_path: &str, server_ocsp_path: &str, server_suite: &str, server_versions: &str, server_session_size: usize, is_server_tickets: bool, server_alpns: &str ) -> IOResult<Self> {
let client_auth_path = if client_auth_path.is_empty() {
None
} else {
Some(Path::new(client_auth_path))
};
let server_certs_path = Path::new(server_certs_path);
let server_key_path = Path::new(server_key_path);
let server_ocsp_path = if server_ocsp_path.is_empty() {
None
} else {
Some(Path::new(server_ocsp_path))
};
let server_suite = if server_suite.is_empty() {
vec![]
} else {
server_suite.split(",").into_iter().map(|suite| {
suite.to_string()
}).collect::<Vec<String>>()
};
let server_versions = if server_versions.is_empty() {
vec![]
} else {
server_versions.split(",").into_iter().map(|version| {
version.to_string()
}).collect::<Vec<String>>()
};
let server_alpns = if server_alpns.is_empty() {
vec![]
} else {
server_alpns.split(",").into_iter().map(|protocol| {
protocol.to_string()
}).collect::<Vec<String>>()
};
match make_server_config(client_auth_path,
is_client_auth,
server_certs_path,
server_key_path,
server_ocsp_path,
server_suite,
server_versions,
server_session_size,
is_server_tickets,
server_alpns) {
Err(e) => Err(e),
Ok(config) => Ok(TlsConfig::Server(config)),
}
}
pub fn new_client() -> Result<Self, String> {
unimplemented!();
}
pub fn is_empty(&self) -> bool {
if let TlsConfig::Empty = self {
return true;
}
false
}
pub fn is_server(&self) -> bool {
if let TlsConfig::Server(_) = self {
return true;
}
false
}
pub fn is_client(&self) -> bool {
if let TlsConfig::Client(_) = self {
return true;
}
false
}
pub fn all_suites(&self) -> Vec<String> {
let mut vec = Vec::new();
for suite in ALL_CIPHER_SUITES {
vec.push(format!("{:?}", suite).to_lowercase());
}
return vec;
}
pub fn all_versions(&self) -> Vec<String> {
vec![format!("{:?}", ProtocolVersion::TLSv1_2).to_lowercase(), format!("{:?}", ProtocolVersion::TLSv1_3).to_lowercase()]
}
}
fn make_server_config(client_auth_path: Option<&Path>,
is_client_auth: bool,
server_certs_path: &Path,
server_key_path: &Path,
server_ocsp_path: Option<&Path>,
server_suite: Vec<String>,
server_versions: Vec<String>,
server_session_size: usize,
is_server_tickets: bool,
_server_alpns: Vec<String>) -> IOResult<Arc<ServerConfig>> {
let client_auth = if let Some(path) = client_auth_path {
match load_certs(path) {
Err(e) => {
return Err(e);
},
Ok(roots) => {
let mut client_auth_roots = RootCertStore::empty();
for root in roots {
if let Err(e) = client_auth_roots.add(&root) {
return Err(Error::new(ErrorKind::Other,
format!("Save client auth root failed, reason: {:?}",
e)));
}
}
if is_client_auth {
AllowAnyAuthenticatedClient::new(client_auth_roots)
} else {
AllowAnyAnonymousOrAuthenticatedClient::new(client_auth_roots)
}
},
}
} else {
NoClientAuth::new()
};
let certs = match load_certs(server_certs_path) {
Err(e) => {
return Err(e);
},
Ok(certs) => {
certs
}
};
let pk = match load_private_key(server_key_path) {
Err(e) => {
return Err(e);
},
Ok(pk) => {
pk
},
};
let ocsp = load_ocsp(&server_ocsp_path)?;
let mut config = match rustls::ServerConfig::builder()
.with_safe_defaults()
.with_client_cert_verifier(client_auth)
.with_single_cert_with_ocsp_and_sct(certs, pk, ocsp, vec![]) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("bad certificates or private key, reason: {:?}",
e)));
},
Ok(cfg) => cfg,
};
config.key_log = Arc::new(rustls::KeyLogFile::new());
if !server_suite.is_empty() {
match select_suites(&server_suite) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Select suites failed, reason: {:?}",
e)));
},
Ok(suites) => {
},
}
}
if !server_versions.is_empty() {
match select_versions(&server_versions) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Select version failed, reason: {:?}",
e)));
},
Ok(versions) => {
},
}
}
if server_session_size != 0 {
}
if is_server_tickets {
}
Ok(Arc::new(config))
}
fn load_certs<P: AsRef<Path>>(file_path: P) -> IOResult<Vec<Certificate>> {
let certfile = match File::open(&file_path) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Open certs failed, path: {:?}, reason: {:?}",
file_path.as_ref(), e)));
},
Ok(file) => {
file
},
};
let mut reader = BufReader::new(certfile);
match rustls_pemfile::certs(&mut reader) {
Err(e) => {
Err(Error::new(ErrorKind::Other,
format!("Load certs failed, path: {:?}, reason: {:?}",
file_path.as_ref(), e)))
},
Ok(certs) => {
Ok(certs.iter()
.map(|v| {
Certificate(v.clone())
})
.collect())
},
}
}
fn load_private_key<P: AsRef<Path>>(file_path: P) -> IOResult<PrivateKey> {
let keyfile = match File::open(&file_path) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Open private key failed, path: {:?}, reason: {:?}",
file_path.as_ref(),
e)));
},
Ok(file) => {
file
},
};
let mut reader = BufReader::new(keyfile);
loop {
match rustls_pemfile::read_one(&mut reader) {
Ok(Some(rustls_pemfile::Item::RSAKey(key))) => return Ok(PrivateKey(key)),
Ok(Some(rustls_pemfile::Item::PKCS8Key(key))) => return Ok(PrivateKey(key)),
Ok(Some(rustls_pemfile::Item::ECKey(key))) => return Ok(PrivateKey(key)),
Ok(None) => break,
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Load private key failed, path: {:?}, reason: {:?}",
file_path.as_ref(),
e)));
},
_ => {
return Err(Error::new(ErrorKind::Other,
format!("Load private key failed, path: {:?}, reason: cannot parse private key .pem file",
file_path.as_ref())))
},
}
}
Err(Error::new(ErrorKind::Other,
format!("no keys found in {:?} (encrypted keys not supported)",
file_path.as_ref())))
}
fn load_ocsp<P: AsRef<Path>>(file_path: &Option<P>) -> IOResult<Vec<u8>> {
let mut ret = Vec::new();
if let Some(path) = &file_path {
match File::open(path) {
Err(e) => {
return Err(Error::new(ErrorKind::Other,
format!("Open ocsp failed, path: {:?}, reason: {:?}",
path.as_ref(),
e)));
},
Ok(mut file) => {
file.read_to_end(&mut ret).unwrap();
},
}
}
Ok(ret)
}
fn select_suites(suites: &[String]) -> Result<Vec<&'static rustls::SupportedCipherSuite>, String> {
let mut result = Vec::new();
for cs_name in suites {
if let Some(suite) = find_suite(cs_name) {
result.push(suite);
} else {
return Err(format!("Cannot select ciphersuite '{}'", cs_name));
}
}
Ok(result)
}
fn find_suite(cs_name: &str) -> Option<&'static rustls::SupportedCipherSuite> {
for suite in rustls::ALL_CIPHER_SUITES {
let sname = format!("{:?}", suite.suite()).to_lowercase();
if sname == cs_name.to_string().to_lowercase() {
return Some(suite);
}
}
None
}
fn select_versions(versions: &[String]) -> Result<Vec<rustls::ProtocolVersion>, String> {
let mut result = Vec::new();
for vname in versions {
let version = match vname.as_ref() {
"1.2" => ProtocolVersion::TLSv1_2,
"1.3" => ProtocolVersion::TLSv1_3,
_ => {
return Err(format!("Cannot select tls version '{}', valid are '1.2' and '1.3'", vname));
},
};
result.push(version);
}
Ok(result)
}
pub struct ContextHandle<T: 'static>(Option<Arc<T>>);
unsafe impl<T: 'static> Send for ContextHandle<T> {}
impl<T: 'static> Drop for ContextHandle<T> {
fn drop(&mut self) {
if let Some(shared) = self.0.take() {
Arc::into_raw(shared);
}
}
}
impl<T: 'static> ContextHandle<T> {
pub fn as_ref(&self) -> &T {
self.0.as_ref().unwrap().as_ref()
}
pub fn as_mut(&mut self) -> Option<&mut T> {
if let Some(shared) = self.0.as_mut() {
return Arc::get_mut(shared);
}
None
}
}
pub struct SocketContext {
inner: *const (), }
unsafe impl Send for SocketContext {}
impl SocketContext {
pub fn empty() -> Self {
SocketContext {
inner: ptr::null(),
}
}
pub fn is_empty(&self) -> bool {
self.inner.is_null()
}
pub fn get<T>(&self) -> Option<ContextHandle<T>> {
if self.is_empty() {
return None;
}
Some(unsafe { ContextHandle(Some(Arc::from_raw(self.inner as *const T))) })
}
pub fn set<T>(&mut self, context: T) -> bool {
if !self.is_empty() {
return false;
}
self.inner = Arc::into_raw(Arc::new(context)) as *const T as *const ();
true
}
pub fn remove<T>(&mut self) -> Result<Option<T>, &str> {
if self.is_empty() {
return Ok(None);
}
let inner = unsafe { Arc::from_raw(self.inner as *const T) };
if Arc::strong_count(&inner) > 1 {
Arc::into_raw(inner); Err("Remove context failed, reason: context shared exist")
} else {
match Arc::try_unwrap(inner) {
Err(inner) => {
Arc::into_raw(inner); Err("Remove context failed, reason: invalid shared")
},
Ok(context) => {
self.inner = ptr::null();
Ok(Some(context))
},
}
}
}
}
pub struct SharedStream<S: Stream>(Arc<RefCell<S>>);
unsafe impl<S: Stream> Send for SharedStream<S> {}
impl<S: Stream> Clone for SharedStream<S> {
fn clone(&self) -> Self {
SharedStream(self.0.clone())
}
}
impl<S: Stream> SharedStream<S> {
pub fn new(s: S) -> Self {
SharedStream(Arc::new(RefCell::new(s)))
}
pub fn with_inner(inner: &Arc<RefCell<S>>) -> Self {
SharedStream(inner.clone())
}
pub fn inner_ref(&self) -> &Arc<RefCell<S>> {
&self.0
}
#[inline]
pub fn borrow<'a>(&'a self) -> SharedRef<'a, S> {
SharedRef(self.0.borrow())
}
#[inline]
pub fn borrow_mut<'a>(&'a self) -> SharedRefMut<'a, S> {
SharedRefMut(self.0.borrow_mut())
}
#[inline]
pub fn as_ptr(&self) -> *mut S {
self.0.as_ptr()
}
}
pub struct SharedRef<'a, S: Stream>(Ref<'a, S>);
unsafe impl<'a, S: Stream> Send for SharedRef<'a, S> {}
impl<'a, S: Stream> Deref for SharedRef<'a, S> {
type Target = S;
#[inline]
fn deref(&self) -> &Self::Target {
&*self.0
}
}
pub struct SharedRefMut<'a, S: Stream>(RefMut<'a, S>);
unsafe impl<'a, S: Stream> Send for SharedRefMut<'a, S> {}
impl<'a, S: Stream> Deref for SharedRefMut<'a, S> {
type Target = S;
#[inline]
fn deref(&self) -> &Self::Target {
&*self.0
}
}
impl<'a, S: Stream> DerefMut for SharedRefMut<'a, S> {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
&mut *self.0
}
}
pub struct Hibernate<S: Socket>(Arc<InnerHibernate<S>>);
impl<S: Socket> Clone for Hibernate<S> {
fn clone(&self) -> Self {
Hibernate(self.0.clone())
}
}
impl<S: Socket> Future for Hibernate<S> {
type Output = IOResult<()>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut locked = self.0.waker.lock();
let result = self
.0
.result
.0
.borrow_mut()
.take();
if let Some(result) = result {
if let Err(e) = self
.0
.handle
.reregister_interest(self.0.ready.clone()) {
Poll::Ready(Err(e))
} else {
self.0.handle.run_hibernated_tasks(); Poll::Ready(result)
}
} else {
if let Err(e) = self
.0
.handle
.reregister_interest(self.0.ready.clone()) {
Poll::Ready(Err(e))
} else {
if self.0.handle.set_hibernate(self.clone()) {
*locked = Some(cx.waker().clone());
} else {
self
.0
.handle
.set_hibernate_wakers(cx.waker().clone());
self.0.waker_status.store(1, Ordering::Relaxed); }
Poll::Pending
}
}
}
}
impl<S: Socket> Hibernate<S> {
pub fn new(handle: SocketHandle<S>,
ready: Ready) -> Self {
let inner = InnerHibernate {
handle,
ready,
waker: SpinLock::new(None),
waker_status: AtomicU8::new(0),
result: AsyncWaitResult(Arc::new(RefCell::new(None))),
};
Hibernate(Arc::new(inner))
}
pub(crate) fn wakeup(&self, result: IOResult<()>) -> bool {
let mut locked = self
.0
.waker
.lock();
if let Some(waker) = locked.take() {
if self.0.waker_status.load(Ordering::Relaxed) > 0 {
self.0.waker_status.store(0, Ordering::Relaxed); waker.wake();
false
} else {
*self.0.result.0.borrow_mut() = Some(result);
waker.wake();
true
}
} else {
if self.0.result.0.borrow().is_none() {
false
} else {
true
}
}
}
}
struct InnerHibernate<S: Socket> {
handle: SocketHandle<S>, ready: Ready, waker: SpinLock<Option<Waker>>, waker_status: AtomicU8, result: AsyncWaitResult<()>, }
#[derive(Debug, Clone)]
pub enum Ready {
Empty, Readable, Writable, OnlyRead, OnlyWrite, ReadWrite, }