#![deny(missing_docs)]
#![doc(html_root_url = "https://docs.rs/tokio-tls/0.1")]
#[cfg_attr(feature = "tokio-proto", macro_use)]
extern crate futures;
extern crate tls_api;
#[macro_use]
extern crate tokio_io;
use std::fmt;
use std::io::{self, Read, Write};
use futures::{Poll, Future, Async};
use tls_api::{HandshakeError, Error, TlsConnector, TlsAcceptor};
use tokio_io::{AsyncRead, AsyncWrite};
pub mod proto;
#[derive(Debug)]
pub struct TlsStream<S> {
inner: tls_api::TlsStream<S>,
}
pub struct ConnectAsync<S> {
inner: MidHandshake<S>,
}
pub struct AcceptAsync<S> {
inner: MidHandshake<S>,
}
struct MidHandshake<S> {
inner: Option<Result<tls_api::TlsStream<S>, HandshakeError<S>>>,
}
impl<S> TlsStream<S> {
pub fn get_ref(&self) -> &tls_api::TlsStream<S> {
&self.inner
}
pub fn get_mut(&mut self) -> &mut tls_api::TlsStream<S> {
&mut self.inner
}
}
impl<S: Read + Write> Read for TlsStream<S> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.inner.read(buf)
}
}
impl<S: Read + Write> Write for TlsStream<S> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.inner.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.inner.flush()
}
}
impl<S: AsyncRead + AsyncWrite> AsyncRead for TlsStream<S> {
}
impl<S: AsyncRead + AsyncWrite + 'static> AsyncWrite for TlsStream<S> {
fn shutdown(&mut self) -> Poll<(), io::Error> {
try_nb!(self.inner.shutdown());
self.inner.get_mut().shutdown()
}
}
pub fn connect_async<C, S>(connector: &C, domain: &str, stream: S) -> ConnectAsync<S>
where
S : io::Read + io::Write + fmt::Debug + Send + Sync + 'static,
C : TlsConnector,
{
ConnectAsync {
inner: MidHandshake {
inner: Some(connector.connect(domain, stream)),
},
}
}
pub fn accept_async<A, S>(acceptor: &A, stream: S) -> AcceptAsync<S>
where
S : io::Read + io::Write + fmt::Debug + Send + Sync + 'static,
A : TlsAcceptor,
{
AcceptAsync {
inner: MidHandshake {
inner: Some(acceptor.accept(stream)),
},
}
}
impl<S: Read + Write + 'static> Future for ConnectAsync<S> {
type Item = TlsStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<TlsStream<S>, Error> {
self.inner.poll()
}
}
impl<S: Read + Write + 'static> Future for AcceptAsync<S> {
type Item = TlsStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<TlsStream<S>, Error> {
self.inner.poll()
}
}
impl<S: Read + Write + 'static> Future for MidHandshake<S> {
type Item = TlsStream<S>;
type Error = Error;
fn poll(&mut self) -> Poll<TlsStream<S>, Error> {
match self.inner.take().expect("cannot poll MidHandshake twice") {
Ok(stream) => Ok(TlsStream { inner: stream }.into()),
Err(HandshakeError::Failure(e)) => Err(e),
Err(HandshakeError::Interrupted(s)) => {
match s.handshake() {
Ok(stream) => Ok(TlsStream { inner: stream }.into()),
Err(HandshakeError::Failure(e)) => Err(e),
Err(HandshakeError::Interrupted(s)) => {
self.inner = Some(Err(HandshakeError::Interrupted(s)));
Ok(Async::NotReady)
}
}
}
}
}
}